From a95f4bdf5c02c71523131e80112ba3d66096cf3e Mon Sep 17 00:00:00 2001 From: leonzzhu Date: Fri, 7 Aug 2026 16:22:43 +0800 Subject: [PATCH 001/483] [Store] Fix stale per-segment metric labels after segment unmount (#3301) MasterMetricManager uses yalantinglibs dynamic_gauge_1t for per-segment metrics (mem_allocated_size_per_segment_, mem_total_capacity_per_segment_, nof_allocated_size_per_segment_, nof_total_capacity_per_segment_). When a segment is unmounted via CommitUnmountSegment, dec_total_mem_capacity() and dec_allocated_nof_size() only decrement the gauge value to 0 but do not remove the label entry from the gauge's internal map. This causes stale 0-value entries to persist indefinitely in Prometheus output. The problem is especially visible after a master restart with snapshot restore: the restored segments carry old client IDs, the reaper eventually expires those clients and calls CommitUnmountSegment, which decrements capacity to 0 but leaves the label behind. After clients remount with new IDs, the old segment names linger as capacity=0 entries. Fix: add remove_segment_metrics() and remove_nof_segment_metrics() that call remove_label_value() on the per-segment gauges, and invoke them in: - ScopedSegmentAccess::CommitUnmountSegment (memory segments) - ScopedNoFSegmentAccess::CommitUnmountSegment (NoF segments) - SegmentManager::releaseCapacityMetrics() (HA teardown) - ~MasterService() standby allocated-size cleanup (HA teardown) Signed-off-by: leonzzhu --- mooncake-store/include/master_metric_manager.h | 7 +++++++ mooncake-store/src/master_metric_manager.cpp | 11 +++++++++++ mooncake-store/src/master_service.cpp | 1 + mooncake-store/src/segment.cpp | 6 ++++++ 4 files changed, 25 insertions(+) diff --git a/mooncake-store/include/master_metric_manager.h b/mooncake-store/include/master_metric_manager.h index 07e565d542..c98c43b177 100644 --- a/mooncake-store/include/master_metric_manager.h +++ b/mooncake-store/include/master_metric_manager.h @@ -92,6 +92,11 @@ class MasterMetricManager { void reset_segment_total_mem_capacity(const std::string& segment); int64_t get_segment_allocated_mem_size(const std::string& segment); int64_t get_segment_total_mem_capacity(const std::string& segment); + // Remove all per-segment metric labels for the given segment. + // Called when a segment is unmounted to prevent stale 0-value entries + // from persisting in Prometheus output (e.g. after snapshot restore + // followed by client expiry / reaper cleanup). + void remove_segment_metrics(const std::string& segment); // NoF segment Metrics void inc_allocated_nof_size(int64_t val = 1); @@ -103,6 +108,8 @@ class MasterMetricManager { double get_segment_nof_used_ratio(const std::string& segment); int64_t get_segment_allocated_nof_size(const std::string& segment); int64_t get_segment_total_nof_capacity(const std::string& segment); + // Remove all per-segment NoF metric labels for the given segment. + void remove_nof_segment_metrics(const std::string& segment); // File Storage Metrics void inc_allocated_file_size(int64_t val = 1); diff --git a/mooncake-store/src/master_metric_manager.cpp b/mooncake-store/src/master_metric_manager.cpp index ea52ec91df..8af2cb8597 100644 --- a/mooncake-store/src/master_metric_manager.cpp +++ b/mooncake-store/src/master_metric_manager.cpp @@ -705,6 +705,11 @@ int64_t MasterMetricManager::get_segment_total_mem_capacity( return mem_total_capacity_per_segment_.value({segment}); } +void MasterMetricManager::remove_segment_metrics(const std::string& segment) { + mem_allocated_size_per_segment_.remove_label_value({{"segment", segment}}); + mem_total_capacity_per_segment_.remove_label_value({{"segment", segment}}); +} + double MasterMetricManager::get_segment_mem_used_ratio( const std::string& segment) { double allocated = get_segment_allocated_mem_size(segment); @@ -775,6 +780,12 @@ int64_t MasterMetricManager::get_segment_total_nof_capacity( return nof_total_capacity_per_segment_.value({segment}); } +void MasterMetricManager::remove_nof_segment_metrics( + const std::string& segment) { + nof_allocated_size_per_segment_.remove_label_value({{"segment", segment}}); + nof_total_capacity_per_segment_.remove_label_value({{"segment", segment}}); +} + double MasterMetricManager::get_segment_nof_used_ratio( const std::string& segment) { double allocated = get_segment_allocated_nof_size(segment); diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index e38fabcd78..a201e8cd05 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -603,6 +603,7 @@ MasterService::~MasterService() { for (const auto& [segment, bytes] : standby_accounted_memory_bytes_) { MasterMetricManager::instance().dec_allocated_mem_size( segment, static_cast(bytes)); + MasterMetricManager::instance().remove_segment_metrics(segment); } // Segments still mounted here never went through CommitUnmountSegment; diff --git a/mooncake-store/src/segment.cpp b/mooncake-store/src/segment.cpp index 8f5b1af04e..8ac943cc7a 100644 --- a/mooncake-store/src/segment.cpp +++ b/mooncake-store/src/segment.cpp @@ -463,6 +463,10 @@ ErrorCode ScopedSegmentAccess::CommitUnmountSegment( if (!is_cxl) { MasterMetricManager::instance().dec_total_mem_capacity( segment_name, metrics_dec_capacity); + // Remove per-segment metric labels entirely to avoid stale 0-value + // entries persisting in Prometheus output (e.g. after snapshot + // restore followed by client expiry / reaper cleanup). + MasterMetricManager::instance().remove_segment_metrics(segment_name); } return ErrorCode::OK; @@ -1474,6 +1478,7 @@ ErrorCode ScopedNoFSegmentAccess::CommitUnmountSegment( nof_segment_manager_->mounted_segments_.erase(segment_id); MasterMetricManager::instance().dec_total_nof_capacity( segment_name, metrics_dec_capacity); + MasterMetricManager::instance().remove_nof_segment_metrics(segment_name); return ErrorCode::OK; } @@ -1572,6 +1577,7 @@ void SegmentManager::releaseCapacityMetrics() { } MasterMetricManager::instance().dec_total_mem_capacity(segment.name, segment.size); + MasterMetricManager::instance().remove_segment_metrics(segment.name); } } From 10f91862a75db2c7572a2e7102579f39d45bf31d Mon Sep 17 00:00:00 2001 From: leonzzhu Date: Fri, 7 Aug 2026 18:22:59 +0800 Subject: [PATCH 002/483] [TE] Fix compiler warnings in IBGDA transport (#3312) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fix three categories of compiler warnings in mooncake-transfer-engine: 1. -Wpointer-arith in ibgda/os.h: void* pointer arithmetic in read_all() and write_all() — cast to char* before arithmetic. 2. -Wsign-compare in mlx5gda.cpp: size_t offset variables compared with int literal -1 — use (size_t)-1 to match the variable type. 3. -Wmissing-field-initializers in ibgda_device_transport.cpp: designated initializer for mlx5gda_qp_devctx omitted mutex, bf_offset, wq_head, wq_tail — add explicit zero initializers. Signed-off-by: leonzzhu --- .../include/transport/device/ibgda/os.h | 5 +++-- .../device/ibgda_device_transport.cpp | 4 ++++ .../src/transport/device/mlx5gda.cpp | 20 +++++++++---------- 3 files changed, 17 insertions(+), 12 deletions(-) diff --git a/mooncake-transfer-engine/include/transport/device/ibgda/os.h b/mooncake-transfer-engine/include/transport/device/ibgda/os.h index bf8bfd4b5c..4cfbd64276 100644 --- a/mooncake-transfer-engine/include/transport/device/ibgda/os.h +++ b/mooncake-transfer-engine/include/transport/device/ibgda/os.h @@ -20,7 +20,7 @@ static inline int close_ret(int fd) { static inline ssize_t read_all(int fd, void *buf, size_t len) { size_t total = 0; while (total < len) { - ssize_t n = read(fd, buf + total, len - total); + ssize_t n = read(fd, static_cast(buf) + total, len - total); if (n == -1) { return -1; } @@ -35,7 +35,8 @@ static inline ssize_t read_all(int fd, void *buf, size_t len) { static inline ssize_t write_all(int fd, const void *buf, size_t len) { size_t total = 0; while (total < len) { - ssize_t n = write(fd, buf + total, len - total); + ssize_t n = + write(fd, static_cast(buf) + total, len - total); if (n == -1) { return -1; } diff --git a/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp b/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp index 18533b0fee..7397ede7d5 100644 --- a/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp +++ b/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp @@ -291,6 +291,7 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { mlx5gda_qp_devctx devctx{ .qpn = qp->qpn, .wqeid_mask = qp->num_wqebb - 1, + .mutex = 0, .wq = split_regions ? reinterpret_cast(qp->wq) : reinterpret_cast( static_cast(ctrl_buf_dev_) + @@ -306,6 +307,9 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { static_cast(ctrl_buf_dev_) + qp->dbr_offset), .bf = static_cast(qp->uar->reg_addr), + .bf_offset = 0, + .wq_head = 0, + .wq_tail = 0, }; cudaMemcpy( static_cast(qp_devctxs_) + i * sizeof(mlx5gda_qp_devctx), diff --git a/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp b/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp index 37b22c0853..4e85aaf489 100644 --- a/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp +++ b/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp @@ -199,13 +199,13 @@ struct mlx5gda_cq* mlx5gda_create_cq( cq_offset = memheap_aligned_alloc(ctrl_buf_heap, num_cqe * sizeof(struct mlx5_cqe64), (size_t)1 << MLX5_ADAPTER_PAGE_SHIFT); - if (cq_offset == -1) { + if (cq_offset == (size_t)-1) { perror("Failed to allocate CQ memory"); goto fail; } dbr_offset = memheap_alloc(ctrl_buf_heap, sizeof(struct mlx5gda_cq_dbr)); - if (dbr_offset == -1) { + if (dbr_offset == (size_t)-1) { perror("Failed to allocate CQ DBR memory"); goto fail; } @@ -287,8 +287,8 @@ struct mlx5gda_cq* mlx5gda_create_cq( free(cq); } if (!split_regions) { - if (cq_offset != -1) memheap_free(ctrl_buf_heap, cq_offset); - if (dbr_offset != -1) memheap_free(ctrl_buf_heap, dbr_offset); + if (cq_offset != (size_t)-1) memheap_free(ctrl_buf_heap, cq_offset); + if (dbr_offset != (size_t)-1) memheap_free(ctrl_buf_heap, dbr_offset); } errno = saved_errno; return NULL; @@ -418,14 +418,14 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( wq_offset = memheap_aligned_alloc( ctrl_buf_heap, num_wqebb * sizeof(struct mlx5gda_wqebb), (size_t)1 << MLX5_ADAPTER_PAGE_SHIFT); - if (wq_offset == -1) { + if (wq_offset == (size_t)-1) { perror("Failed to allocate WQ memory"); goto fail; } dbr_offset = memheap_alloc(ctrl_buf_heap, sizeof(struct mlx5gda_wq_dbr)); - if (dbr_offset == -1) { + if (dbr_offset == (size_t)-1) { perror("Failed to allocate DBR memory"); goto fail; } @@ -518,10 +518,10 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( if (qp) { free(qp); } - if (!split_regions && wq_offset != -1) { + if (!split_regions && wq_offset != (size_t)-1) { memheap_free(ctrl_buf_heap, wq_offset); } - if (!split_regions && dbr_offset != -1) { + if (!split_regions && dbr_offset != (size_t)-1) { memheap_free(ctrl_buf_heap, dbr_offset); } errno = saved_errno; @@ -542,10 +542,10 @@ void mlx5gda_destroy_qp(struct memheap* ctrl_buf_heap, struct mlx5gda_qp* qp) { release_control_region(&qp->region_allocator, &qp->dbr_region); release_control_region(&qp->region_allocator, &qp->wq_region); } else { - if (qp->wq_offset != -1) { + if (qp->wq_offset != (size_t)-1) { memheap_free(ctrl_buf_heap, qp->wq_offset); } - if (qp->dbr_offset != -1) { + if (qp->dbr_offset != (size_t)-1) { memheap_free(ctrl_buf_heap, qp->dbr_offset); } } From 6edc24c5defe181059cbc630a21c60f4aff69e1c Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Sat, 8 Aug 2026 01:39:08 +0800 Subject: [PATCH 003/483] [Wheel] Support jagged NestedTensor transfer (#3336) * [wheel] support jagged NestedTensor transfer * [wheel] optimize nested tensor deserialization --- .../mooncake/structured_object_store.py | 118 ++++++++++++++++++ .../tests/test_structured_object_store.py | 49 ++++++++ 2 files changed, 167 insertions(+) diff --git a/mooncake-wheel/mooncake/structured_object_store.py b/mooncake-wheel/mooncake/structured_object_store.py index 9b6a7fbda9..d2a13e5223 100644 --- a/mooncake-wheel/mooncake/structured_object_store.py +++ b/mooncake-wheel/mooncake/structured_object_store.py @@ -871,6 +871,12 @@ def _read_dataproto_member_indices( member, payload_spec, field_spec, indices, destination ) if encoding == "torch_tensor": + if field_spec.get("nested"): + result = self.materialize_into( + self.read_spec(stage_ref).select_members([member]), + None if destination is None else {member: destination}, + ) + return _select_nested_tensor_rows(result.objects[member], indices) return self._read_torch_tensor_member_indices( member, payload_spec, field_spec, indices, destination ) @@ -3178,6 +3184,22 @@ def _read_torch_tensor_member( member_slice: StructuredMemberSlice | None, destination: Any, ) -> Any: + if field_spec.get("nested"): + if destination is not None: + raise ValueError( + f"structured nested tensor member {name} does not support destinations" + ) + value = _deserialize_nested_tensor_payload( + self._bundle_store.read_payload(payload_spec), field_spec + ) + if member_slice is None: + return value + if member_slice.axis != 0: + raise ValueError( + "structured nested tensor slicing currently supports axis=0 only" + ) + start, end, step = _normalized_member_slice(member_slice, len(value)) + return _select_nested_tensor_rows(value, range(start, end, step)) if member_slice is not None: return self._read_sliced_torch_tensor_member( name, payload_spec, field_spec, member_slice, destination @@ -4896,6 +4918,25 @@ def _normalized_member_slice( ) +def _select_nested_tensor_rows(value: Any, indices: Sequence[int]) -> Any: + rows = value.unbind() + selected = [rows[int(index)] for index in indices] + if selected: + result = _torch.nested.as_nested_tensor(selected, layout=value.layout) + else: + offsets = _torch.zeros( + 1, + dtype=value.offsets().dtype, + device=value.offsets().device, + ) + result = _torch.nested.nested_tensor_from_jagged( + value.values()[:0], offsets=offsets + ) + if hasattr(value, "_ragged_idx"): + result._ragged_idx = value._ragged_idx + return result + + def _bytes_view(value: Any, name: str) -> memoryview: try: view = memoryview(value) @@ -5128,6 +5169,10 @@ def _encode_structured_field(value: Any) -> tuple[dict[str, Any], Any]: def _encode_torch_tensor_field(value: Any) -> tuple[dict[str, Any], Any]: + if value.is_nested: + if value.layout != _torch.jagged: + raise ValueError("structured nested tensor fields require torch.jagged layout") + return _encode_nested_tensor_field(value) return { "encoding": "torch_tensor", "dtype": str(value.dtype), @@ -5413,6 +5458,79 @@ def _deserialize_torch_save_payload(payload: bytes) -> Any: return _torch.load(io.BytesIO(payload), weights_only=True) +def _encode_nested_tensor_field(value: Any) -> tuple[dict[str, Any], Any]: + field_spec = { + "encoding": "torch_tensor", + "dtype": str(value.dtype), + "nested": True, + "format": "torch_save", + } + values = value.values() + offsets = value.offsets() + lengths = value.lengths() + tensors = [values, offsets] + if lengths is not None: + tensors.append(lengths) + + if _has_tensor_codec_helpers() and all( + tensor.device.type == "cpu" for tensor in tensors + ): + parts = [ + _tensor_payload_bytes(_TensorPayload(tensor=tensor))[0] + for tensor in tensors + ] + field_spec.update( + { + "format": "tensor_parts", + "part_bytes": [len(part) for part in parts], + "has_lengths": lengths is not None, + "ragged_idx": int(getattr(value, "_ragged_idx", 1)), + } + ) + return field_spec, _MultiBufferPayload( + buffers=tuple(memoryview(part) for part in parts) + ) + + return field_spec, memoryview(_torch_save_payload_bytes(value)) + + +def _deserialize_nested_tensor_payload( + payload: bytes, field_spec: Mapping[str, Any] +) -> Any: + payload_format = field_spec.get("format", "torch_save") + if payload_format == "torch_save": + return _deserialize_torch_save_payload(payload) + if payload_format != "tensor_parts": + raise ValueError(f"unsupported nested tensor payload format: {payload_format}") + + part_bytes = field_spec.get("part_bytes") + expected_parts = 3 if field_spec.get("has_lengths") else 2 + if not isinstance(part_bytes, list) or len(part_bytes) != expected_parts: + raise ValueError("nested tensor payload has invalid part metadata") + + tensors = [] + offset = 0 + for part_size in part_bytes: + if ( + not isinstance(part_size, int) or isinstance(part_size, bool) or part_size <= 0 + ): + raise ValueError("nested tensor payload has invalid part size") + end = offset + int(part_size) + if end > len(payload): + raise ValueError("nested tensor payload is truncated") + tensors.append(_deserialize_tensor_payload(payload[offset:end])) + offset = end + if offset != len(payload): + raise ValueError("nested tensor payload has trailing bytes") + + return _torch.nested.nested_tensor_from_jagged( + tensors[0], + offsets=tensors[1], + lengths=tensors[2] if expected_parts == 3 else None, + jagged_dim=int(field_spec.get("ragged_idx", 1)), + ) + + def _slice_tensor_metadata( metadata: bytes, shape: Sequence[int], data_bytes: int ) -> bytes: diff --git a/mooncake-wheel/tests/test_structured_object_store.py b/mooncake-wheel/tests/test_structured_object_store.py index ad0060d1d0..13131d290d 100644 --- a/mooncake-wheel/tests/test_structured_object_store.py +++ b/mooncake-wheel/tests/test_structured_object_store.py @@ -3041,6 +3041,55 @@ def test_dataproto_helper_ragged_tensor_non_tensor_roundtrip() -> None: transfer.put_dataproto(SimpleDataProto(non_tensor_batch={"ragged": mixed})) +def test_dataproto_helper_jagged_nested_batch_tensor_roundtrip() -> None: + torch = pytest.importorskip("torch") + _store, transfer = make_transfer() + rows = [ + torch.arange(2, dtype=torch.int64), + torch.arange(10, 14, dtype=torch.int64), + torch.arange(20, 21, dtype=torch.int64), + ] + nested = torch.nested.as_nested_tensor(rows, layout=torch.jagged) + + ref = transfer.put_dataproto(SimpleDataProto(batch={"tokens": nested})) + view = transfer.dataproto_manifest_view(ref) + + field_spec = view["batch_fields"]["tokens"]["spec"] + assert field_spec["encoding"] == "torch_tensor" + assert field_spec["dtype"] == "torch.int64" + assert field_spec["nested"] is True + assert field_spec["format"] in {"tensor_parts", "torch_save"} + if field_spec["format"] == "tensor_parts": + assert len(field_spec["part_bytes"]) == 2 + + for selection, expected_rows in ( + (None, rows), + (slice(1, 3), rows[1:3]), + ([2, 0], [rows[2], rows[0]]), + ([], []), + ): + result = transfer.get_dataproto(ref, rows=selection)["batch"]["tokens"] + assert result.is_nested + assert result.layout == torch.jagged + assert len(result) == len(expected_rows) + for actual, expected in zip(result.unbind(), expected_rows): + assert torch.equal(actual, expected) + + matrix_rows = [ + torch.arange(6, dtype=torch.int64).reshape(2, 3), + torch.arange(10, dtype=torch.int64).reshape(2, 5), + ] + matrix = torch.nested.as_nested_tensor(matrix_rows, layout=torch.jagged) + matrix_ref = transfer.put_dataproto(SimpleDataProto(batch={"position_ids": matrix})) + matrix_result = transfer.get_dataproto(matrix_ref)["batch"]["position_ids"] + + assert matrix_result.is_nested + assert matrix_result.layout == torch.jagged + assert matrix_result._ragged_idx == 2 + for actual, expected in zip(matrix_result.unbind(), matrix_rows): + assert torch.equal(actual, expected) + + def _assert_tensor_object_equal(actual, expected) -> None: torch = pytest.importorskip("torch") if expected is None: From 41bfbebe4d93ae58a30cb0aa0ad1b61e275a9ca4 Mon Sep 17 00:00:00 2001 From: Zhanhao Cao Date: Sat, 8 Aug 2026 10:54:40 +0800 Subject: [PATCH 004/483] [PG] Support inplace rejoin and fix p2p recovery (#3323) --- mooncake-pg/include/mooncake_worker.cuh | 6 +- mooncake-pg/include/p2p_proxy.h | 21 +- mooncake-pg/src/mooncake_communicator.cpp | 40 ++- mooncake-pg/src/p2p_proxy.cpp | 115 ++++--- mooncake-pg/tests/test_pg_elastic.py | 303 ++++++++++++++++++ mooncake-pg/tests/test_pg_p2p.py | 9 + mooncake-pg/tests/transfer_fault_injection.py | 231 +++++++++++++ mooncake-pg/tests/transfer_fault_preload.cpp | 189 +++++++++++ 8 files changed, 820 insertions(+), 94 deletions(-) create mode 100644 mooncake-pg/tests/transfer_fault_injection.py create mode 100644 mooncake-pg/tests/transfer_fault_preload.cpp diff --git a/mooncake-pg/include/mooncake_worker.cuh b/mooncake-pg/include/mooncake_worker.cuh index 6efdc58389..a92842d2f2 100644 --- a/mooncake-pg/include/mooncake_worker.cuh +++ b/mooncake-pg/include/mooncake_worker.cuh @@ -31,9 +31,13 @@ class MooncakeCommunicator; // Isolated --(Active view)--> Normal // // Joining member: -// Isolated --(joinGroup: drain preparation collectives)--> Quiescing +// Isolated ---------(joinGroup: drain operations)--------> Quiescing // -----------------(Active view)----------------> Normal // +// In-place rejoining member: +// Normal + inactive --(joinGroup: drain operations)------> Quiescing +// ---------(Active view)---------------> Normal +// // Isolated admits local-only collectives with an {self} active ranks mask. // Quiescing rejects new collectives while waiting for activation. Normal uses // the Coordinator's committed membership as active ranks. These are local diff --git a/mooncake-pg/include/p2p_proxy.h b/mooncake-pg/include/p2p_proxy.h index 6b17d126cf..6e823a81a4 100644 --- a/mooncake-pg/include/p2p_proxy.h +++ b/mooncake-pg/include/p2p_proxy.h @@ -335,15 +335,6 @@ class P2PProxy { */ void abandonResources(); - // Epoch for fault recovery. All control slots carry this - // value so that stale messages from before a Reset can be detected. - uint32_t getEpoch(int peer_rank) const { - return peer_epoch_[peer_rank].load(std::memory_order_acquire); - } - void setEpoch(int peer_rank, uint32_t epoch) { - peer_epoch_[peer_rank].store(epoch, std::memory_order_release); - } - private: // Sender-side per-chunk state machine. enum class SendTaskState { @@ -375,7 +366,7 @@ class P2PProxy { SendTransferTask() = default; SendTransferTask(uint64_t buffer_offset_in, uint32_t chunk_len_in, void* staging_addr_in, uint64_t remote_addr_in, - uint32_t sequence_in, uint32_t epoch_in); + uint32_t sequence_in); SendTaskState state_ = SendTaskState::kCopyIn; uint64_t buffer_offset_ = 0; // Offset inside the user buffer. @@ -383,7 +374,6 @@ class P2PProxy { void* staging_addr_ = nullptr; // Address inside SendPool. uint64_t remote_addr_ = 0; // Address inside REMOTE RecvPool. uint32_t sequence_ = 0; // Sequence number in the control ring. - uint32_t epoch_ = 0; // Epoch for Reset detection. std::optional transfer_batch_id_; // RDMA Write batch id. std::optional ack_batch_id_; // AckSlot write batch id. cudaEvent_t copy_ready_event_ = nullptr; // Signals Copy-In done (GPU). @@ -420,15 +410,13 @@ class P2PProxy { struct RecvTransferTask { RecvTransferTask() = default; RecvTransferTask(uint64_t buffer_offset_in, uint32_t chunk_len_in, - void* local_addr_in, uint32_t sequence_in, - uint32_t epoch_in); + void* local_addr_in, uint32_t sequence_in); RecvTaskState state_ = RecvTaskState::kIssueCredit; uint64_t buffer_offset_ = 0; // Offset inside the user buffer. uint32_t chunk_len_ = 0; // Bytes in this chunk. void* local_addr_ = nullptr; // Address inside local RecvPool. uint32_t sequence_ = 0; // Sequence number in the control ring. - uint32_t epoch_ = 0; // Epoch for Reset detection. std::optional credit_batch_id_; // CreditSlot RDMA Write batch id. cudaEvent_t copy_ready_event_ = @@ -591,11 +579,6 @@ class P2PProxy { std::atomic active_send_tasks_{0}; std::atomic active_recv_tasks_{0}; - // Per-peer epoch for fault recovery. Incremented in resetPeerState - // and performSend/RecvReset so that stale messages from a previous epoch - // can be detected on a per-peer basis. - std::array, kMaxNumRanks> peer_epoch_; - std::array send_peer_lanes_; std::array recv_peer_lanes_; }; diff --git a/mooncake-pg/src/mooncake_communicator.cpp b/mooncake-pg/src/mooncake_communicator.cpp index 647b7361d0..7ad4d4d33a 100644 --- a/mooncake-pg/src/mooncake_communicator.cpp +++ b/mooncake-pg/src/mooncake_communicator.cpp @@ -616,15 +616,16 @@ PGResult MooncakeCommunicator::checkOpState(OpType op) const { "rank " + std::to_string(meta_->globalRank) + " is offline and cannot perform operations"); } - // P2P operations don't require the rank to be active in the group. + PG_VALIDATE_STATE(mode != CollectiveExtensionState::Quiescing, + "rank is quiescing and cannot issue operations"); + const bool is_p2p = op == OpType::Send || op == OpType::Recv; - if (!isValidGroup() && is_p2p) { - return makePGError(PGErrorCode::NotSupported, - "P2P is unavailable for an invalid Mooncake group"); - } - if (!is_p2p) { - PG_VALIDATE_STATE(mode != CollectiveExtensionState::Quiescing, - "rank is quiescing and cannot issue collectives"); + if (is_p2p) { + // P2P operations require an valid group + PG_VALIDATE_STATE(isValidGroup(), + "P2P is unavailable for an invalid group"); + } else { + // Collectives in an invalid group remain local-only PG_VALIDATE_STATE(mode == CollectiveExtensionState::Isolated || meta_->activeRanks[rank_], "rank is not active in this group"); @@ -1433,19 +1434,32 @@ PGResult MooncakeCommunicator::deactivateRanks( PGResult MooncakeCommunicator::joinGroup() { PG_TRY(checkValidGroup("joinGroup")); auto mode = meta_->extensionMode.load(std::memory_order_acquire); - PG_VALIDATE_STATE( - mode == CollectiveExtensionState::Isolated, - "joinGroup may only be called once on an isolated joining " - "communicator"); - // Stop admitting isolated collectives before advertising readiness. + const bool is_initial_join = mode == CollectiveExtensionState::Isolated; + const bool is_inplace_rejoin = + mode == CollectiveExtensionState::Normal && !meta_->activeRanks[rank_]; + PG_VALIDATE_STATE(is_initial_join || is_inplace_rejoin, + "joinGroup requires an isolated or inactive rank"); + // Stop admitting operations before advertising readiness. meta_->extensionMode.store(CollectiveExtensionState::Quiescing, std::memory_order_release); + if (!p2p_proxy_->drainTasks()) { + return makePGError(PGErrorCode::Timeout, + "timed out draining join preparation P2P tasks for " + "rank " + + std::to_string(meta_->globalRank)); + } if (!worker_->drainTasks(meta_.get())) { return makePGError( PGErrorCode::Timeout, "timed out draining join preparation collectives for rank " + std::to_string(meta_->globalRank)); } + // Auto-deactivation removes the old endpoint from the GroupView. The + // process and its registered buffers are still alive, so republish the + // current endpoint under a fresh endpoint epoch before declaring ready. + if (is_inplace_rejoin) { + PG_TRY(agent_.publishLocalEndpoint(buildEndpointMetadata())); + } PG_TRY(agent_.confirmReadyForActivation(meta_->group_id)); // Block until the Coordinator activates this rank in the group. PG_TRY(agent_.waitUntilRankActive(meta_->group_id, meta_->globalRank, diff --git a/mooncake-pg/src/p2p_proxy.cpp b/mooncake-pg/src/p2p_proxy.cpp index 074eb692fe..109819e00c 100644 --- a/mooncake-pg/src/p2p_proxy.cpp +++ b/mooncake-pg/src/p2p_proxy.cpp @@ -218,10 +218,6 @@ void P2PProxy::allocateResources() { resources_.ack_region_[i].reset(); } - for (size_t i = 0; i < kMaxNumRanks; ++i) { - peer_epoch_[i].store(1, std::memory_order_release); - } - for (size_t i = 0; i < kMaxNumRanks; ++i) { send_peer_lanes_[i].pending_send_ops_.clear(); send_peer_lanes_[i].active_send_op_.reset(); @@ -287,12 +283,6 @@ void P2PProxy::resetPeerState(int peer_rank) { "ResetPeerState: peer_rank out of range: ", peer_rank, " size: ", size_); - // Epoch update: - // We bump epoch_ so that any slots still in flight from the - // old session (written by the sender/receiver before it learned about - // the Reset) are recognized as stale and skipped. - peer_epoch_[peer_rank].fetch_add(1, std::memory_order_acq_rel); - // Request reset reset_send_req_[peer_rank].store(true, std::memory_order_release); reset_recv_req_[peer_rank].store(true, std::memory_order_release); @@ -332,22 +322,35 @@ void P2PProxy::cleanupFailedRecvOp(RecvOpContext& op_ctx) { active_recv_tasks_.fetch_sub(1, std::memory_order_release); } -void P2PProxy::handleFailedSendOp(SendOpContext& op_ctx) { - cleanupFailedSendOp(op_ctx); - op_ctx.failed_ranks_hint_[op_ctx.peer_rank_] = 1; +void P2PProxy::reportPeerFailure(int peer_rank) { // Reset P2P session state (epoch, lanes). - resetPeerState(op_ctx.peer_rank_); + resetPeerState(peer_rank); // link event vector is indexed by GlobalRank. - const auto peer_global = meta_->rank_order[op_ctx.peer_rank_]; + const auto peer_global = meta_->rank_order[peer_rank]; const auto target_rank_epoch = meta_->rankEpochs[peer_global]; - if (meta_->communicator) { - LinkEvent event; - event.events.assign(kMaxNumRanks, LinkEvent::EventType::None); - event.target_rank_epochs.assign(kMaxNumRanks, 0); - event.events[peer_global] = LinkEvent::EventType::Failure; - event.target_rank_epochs[peer_global] = target_rank_epoch; - meta_->communicator->getAgent().pushLinkEvent(event); + auto* communicator = meta_->communicator; + if (!communicator) return; + + LinkEvent event; + event.events.assign(kMaxNumRanks, LinkEvent::EventType::None); + event.target_rank_epochs.assign(kMaxNumRanks, 0); + event.events[peer_global] = LinkEvent::EventType::Failure; + event.target_rank_epochs[peer_global] = target_rank_epoch; + communicator->getAgent().pushLinkEvent(event); + + if (meta_->autoSyncOnFailure) { + auto result = communicator->syncAfterFailure(); + PG_ASSERT(result.has_value() && + result.value().status != SyncAfterFailureStatus::Rejected, + "syncAfterFailure failed for rank ", meta_->globalRank); } +} + +void P2PProxy::handleFailedSendOp(SendOpContext& op_ctx) { + cleanupFailedSendOp(op_ctx); + op_ctx.failed_ranks_hint_[op_ctx.peer_rank_] = 1; + reportPeerFailure(op_ctx.peer_rank_); + const auto peer_global = meta_->rank_order[op_ctx.peer_rank_]; op_ctx.completion_->set_value(); LOG(ERROR) << "Rank " << meta_->rank << ": P2P SendOp to peer " << op_ctx.peer_rank_ << " (global=" << peer_global @@ -357,19 +360,8 @@ void P2PProxy::handleFailedSendOp(SendOpContext& op_ctx) { void P2PProxy::handleFailedRecvOp(RecvOpContext& op_ctx) { cleanupFailedRecvOp(op_ctx); op_ctx.failed_ranks_hint_[op_ctx.peer_rank_] = 1; - // Reset P2P session state (epoch, lanes). - resetPeerState(op_ctx.peer_rank_); - // link event vector is indexed by GlobalRank. + reportPeerFailure(op_ctx.peer_rank_); const auto peer_global = meta_->rank_order[op_ctx.peer_rank_]; - const auto target_rank_epoch = meta_->rankEpochs[peer_global]; - if (meta_->communicator) { - LinkEvent event; - event.events.assign(kMaxNumRanks, LinkEvent::EventType::None); - event.target_rank_epochs.assign(kMaxNumRanks, 0); - event.events[peer_global] = LinkEvent::EventType::Failure; - event.target_rank_epochs[peer_global] = target_rank_epoch; - meta_->communicator->getAgent().pushLinkEvent(event); - } op_ctx.completion_->set_value(); LOG(ERROR) << "Rank " << meta_->rank << ": P2P RecvOp from peer " << op_ctx.peer_rank_ << " (global=" << peer_global @@ -518,15 +510,16 @@ void P2PProxy::enqueueRecv(RecvOp op) { if (device_worker_) device_worker_->wakeUpRecv(); } -P2PProxy::SendTransferTask::SendTransferTask( - uint64_t buffer_offset_in, uint32_t chunk_len_in, void* staging_addr_in, - uint64_t remote_addr_in, uint32_t sequence_in, uint32_t epoch_in) +P2PProxy::SendTransferTask::SendTransferTask(uint64_t buffer_offset_in, + uint32_t chunk_len_in, + void* staging_addr_in, + uint64_t remote_addr_in, + uint32_t sequence_in) : buffer_offset_(buffer_offset_in), chunk_len_(chunk_len_in), staging_addr_(staging_addr_in), remote_addr_(remote_addr_in), - sequence_(sequence_in), - epoch_(epoch_in) { + sequence_(sequence_in) { last_update_time_ = std::chrono::steady_clock::now(); } @@ -543,13 +536,11 @@ P2PProxy::SendOpContext::SendOpContext(SendOp&& op_in) P2PProxy::RecvTransferTask::RecvTransferTask(uint64_t buffer_offset_in, uint32_t chunk_len_in, void* local_addr_in, - uint32_t sequence_in, - uint32_t epoch_in) + uint32_t sequence_in) : buffer_offset_(buffer_offset_in), chunk_len_(chunk_len_in), local_addr_(local_addr_in), - sequence_(sequence_in), - epoch_(epoch_in) { + sequence_(sequence_in) { last_update_time_ = std::chrono::steady_clock::now(); } @@ -632,15 +623,15 @@ bool P2PProxy::tryIssueRecvTask(RecvOpContext& op_ctx, RecvPeerLane& lane) { const uint64_t remote_credit_offset = getRemoteCreditSlot(op_ctx.peer_rank_, seq); - const uint32_t curr_epoch = - peer_epoch_[op_ctx.peer_rank_].load(std::memory_order_acquire); + const uint32_t group_epoch = + static_cast(meta_->epoch.load(std::memory_order_acquire)); op_ctx.tasks_.emplace_back(op_ctx.bytes_credited_, chunk_len, local_addr, - seq, curr_epoch); + seq); auto& task = op_ctx.tasks_.back(); auto* credit_staging_buf = getLocalCreditStagingBuf(op_ctx.peer_rank_, seq); credit_staging_buf->publish( - curr_epoch, seq, reinterpret_cast(local_addr), chunk_len); + group_epoch, seq, reinterpret_cast(local_addr), chunk_len); const BatchID batch_id = engine_->allocateBatchID(1); engine_->submitTransfer( batch_id, {TransferRequest{ @@ -768,15 +759,15 @@ P2PProxy::IssueResult P2PProxy::tryIssueSendTask(SendOpContext& op_ctx, return IssueResult::kNoCredit; } - // Step 2 -- Stale packet: the slot carries data from a previous epoch - // (before a Reset). Clear it so the fresh credit can land safely. - const uint32_t curr_epoch = - peer_epoch_[op_ctx.peer_rank_].load(std::memory_order_acquire); - if (slot_epoch != curr_epoch) { + // Step 2 -- Stale packet: the slot belongs to an older group view. Clear it + // so the fresh credit can land safely. + const uint32_t group_epoch = + static_cast(meta_->epoch.load(std::memory_order_acquire)); + if (slot_epoch != group_epoch) { LOG(WARNING) << "[P2PProxy][Send] tryIssueSendTask peer=" << op_ctx.peer_rank_ << " STALE_EPOCH seq=" << seq << " slot.epoch=" << slot_epoch - << " curr_epoch=" << curr_epoch; + << " group_epoch=" << group_epoch; slot.reset(); return IssueResult::kIssued; } @@ -798,7 +789,7 @@ P2PProxy::IssueResult P2PProxy::tryIssueSendTask(SendOpContext& op_ctx, "P2P send got invalid chunk_len in credit slot."); op_ctx.tasks_.emplace_back(op_ctx.bytes_staged_, chunk_len, staging_addr, - recv_addr, seq, slot_epoch); + recv_addr, seq); auto& task = op_ctx.tasks_.back(); slot.reset(); @@ -939,7 +930,9 @@ bool P2PProxy::stepSendAck(SendOpContext& op_ctx, SendTransferTask& task) { if (!task.ack_batch_id_.has_value()) { auto* ack_staging_buf = getLocalAckStagingBuf(op_ctx.peer_rank_, task.sequence_); - ack_staging_buf->publish(task.epoch_, task.sequence_, task.chunk_len_); + const uint32_t group_epoch = + static_cast(meta_->epoch.load(std::memory_order_acquire)); + ack_staging_buf->publish(group_epoch, task.sequence_, task.chunk_len_); const BatchID batch_id = engine_->allocateBatchID(1); engine_->submitTransfer( @@ -1123,17 +1116,17 @@ bool P2PProxy::pollRecvAckSlot(RecvOpContext& op_ctx, RecvPeerLane& lane, return false; } - const uint32_t curr_epoch = - peer_epoch_[op_ctx.peer_rank_].load(std::memory_order_acquire); + const uint32_t group_epoch = + static_cast(meta_->epoch.load(std::memory_order_acquire)); - // Step 2 -- Stale packet: data from a previous epoch (before Reset). - // Clear the slot so the fresh ack can land safely. - if (slot_epoch != curr_epoch) { + // Step 2 -- Stale packet from an older group view. Clear the slot so the + // fresh ack can land safely. + if (slot_epoch != group_epoch) { LOG(WARNING) << "[P2PProxy][Recv] pollRecvAckSlot peer=" << op_ctx.peer_rank_ << " front-seq=" << head_task.sequence_ << " EPOCH_MISMATCH slot.epoch=" << slot_epoch - << " curr_epoch=" << curr_epoch; + << " group_epoch=" << group_epoch; slot.reset(); return true; } diff --git a/mooncake-pg/tests/test_pg_elastic.py b/mooncake-pg/tests/test_pg_elastic.py index 114bdefc94..db7d35e0f2 100644 --- a/mooncake-pg/tests/test_pg_elastic.py +++ b/mooncake-pg/tests/test_pg_elastic.py @@ -15,6 +15,10 @@ require_test_device, wait_until, ) +from transfer_fault_injection import ( + TransferFault, + preload_transfer_fault, +) BROKEN_RANK = 1 @@ -775,6 +779,250 @@ def _replacement_recovery_worker( ctx.record_result({"role": "replacement"}) +def _run_p2p_ping_pong( + ctx: MooncakePGWorkerContext, + device: torch.device, + logical_rank: int, +) -> None: + """Exchange one message in each direction between two logical ranks.""" + if ctx.world_size != 2: + raise AssertionError("recovery P2P test expects exactly two ranks") + + peer = 1 - logical_rank + send_tensor = torch.tensor( + [logical_rank], dtype=torch.int64, device=device + ) + recv_tensor = torch.empty_like(send_tensor) + works = dist.batch_isend_irecv( + [ + dist.P2POp(op=dist.isend, tensor=send_tensor, peer=peer), + dist.P2POp(op=dist.irecv, tensor=recv_tensor, peer=peer), + ] + ) + for work in works: + work.wait() + + if not all(pg.get_local_success(work) for work in works): + raise AssertionError( + f"rank {logical_rank}: P2P with rank {peer} failed" + ) + received = int(recv_tensor.cpu().item()) + if received != peer: + raise AssertionError( + f"rank {logical_rank}: expected {peer}, received {received}" + ) + + +def _recovery_p2p_worker( + ctx: MooncakePGWorkerContext, + broken_exited: mp.Event, + start_recovery: mp.Event, + graceful_group_destroy: bool, +) -> None: + """Replace logical rank 1, then repeat its P2P exchange with rank 0.""" + logical_rank = ctx.rank if ctx.proc_rank < ctx.world_size else BROKEN_RANK + + if ctx.proc_rank < ctx.world_size: + device = ctx.init_group(rank=logical_rank) + backend = ctx.get_backend() + _run_p2p_ping_pong(ctx, device, logical_rank) + + if logical_rank == BROKEN_RANK: + if graceful_group_destroy: + resp = pg.deactivate_ranks(backend, [logical_rank]) + assert resp.status == pg.ProposalStatus.Applied, \ + "graceful self-deactivation should apply, " \ + f"got {resp.status}: {resp.reject_reason}" + + dist.destroy_process_group() + ctx.record_result({"role": "gracefully_removed"}) + return + + ctx.record_result({"role": "broken"}) + broken_exited.set() + os._exit(0) + + if not broken_exited.wait(timeout=120.0): + raise TimeoutError("timed out waiting for departed rank") + + if not graceful_group_destroy: + probe = torch.tensor([logical_rank], device=device) + work = dist.isend(probe, dst=BROKEN_RANK) + work.wait() + if pg.get_local_success(work): + raise AssertionError( + "P2P probe to the departed rank unexpectedly succeeded" + ) + + active_ranks = pg.get_active_ranks(backend).cpu().tolist() + assert active_ranks == [1, 0], ( + "new group view should be applied before p2p completes, " + f"got active_ranks={active_ranks}" + ) + + start_recovery.set() + + wait_until( + lambda: pg.get_peer_state(backend, [BROKEN_RANK])[0], + timeout_s=30.0, + poll_interval_s=0.05, + description=f"rank {logical_rank} waiting for replacement", + ) + resp = pg.recover_ranks(backend, [BROKEN_RANK]) + assert resp.status == pg.ProposalStatus.Applied, \ + f"rank {logical_rank}: recover_ranks should apply, got {resp.status}" + + role = "survivor" + else: + if not start_recovery.wait(timeout=60.0): + raise TimeoutError("timed out waiting to start replacement") + + device = ctx.init_group(rank=logical_rank, is_extension=True) + backend = ctx.get_backend() + pg.join_group(backend) + role = "replacement" + + _run_p2p_ping_pong(ctx, device, logical_rank) + ctx.record_result({"role": role}) + + +def _run_rejoin_operation( + ctx: MooncakePGWorkerContext, + device: torch.device, + operation: str, + *, + verify: bool, +): + if operation == "collective": + tensor = torch.tensor( + [ctx.rank + 1], dtype=torch.int32, device=device + ) + works = [ + dist.all_reduce(tensor, op=dist.ReduceOp.SUM, async_op=True) + ] + for work in works: + work.wait() + + if verify: + if not pg.get_local_success(works[0]): + raise AssertionError("collective failed locally") + expected = ctx.world_size * (ctx.world_size + 1) // 2 + actual = int(tensor.cpu().item()) + if actual != expected: + raise AssertionError( + f"collective expected {expected}, got {actual}" + ) + return works + + if operation == "p2p": + peer = 1 - ctx.rank + send_tensor = torch.tensor( + [ctx.rank], dtype=torch.int64, device=device + ) + recv_tensor = torch.empty_like(send_tensor) + works = dist.batch_isend_irecv( + [ + dist.P2POp(op=dist.isend, tensor=send_tensor, peer=peer), + dist.P2POp(op=dist.irecv, tensor=recv_tensor, peer=peer), + ] + ) + for work in works: + work.wait() + + if verify: + if not all(pg.get_local_success(work) for work in works): + raise AssertionError(f"P2P with rank {peer} failed locally") + actual = int(recv_tensor.cpu().item()) + if actual != peer: + raise AssertionError(f"expected {peer}, received {actual}") + return works + + raise ValueError(f"unsupported in-place rejoin operation: {operation}") + + +def _run_fault_round( + ctx: MooncakePGWorkerContext, + device: torch.device, + backend, + operation: str, + fault_library: str, + fault_started: mp.Event, +) -> None: + fault = TransferFault(fault_library, local_rank=ctx.rank) + with fault.failing_link(BROKEN_RANK, 0): + if ctx.rank != BROKEN_RANK: + if not fault_started.wait(timeout=30.0): + raise TimeoutError("timed out waiting for fault injection") + _run_rejoin_operation(ctx, device, operation, verify=False) + return + + fault_started.set() + works = _run_rejoin_operation( + ctx, device, operation, verify=False + ) + + if all(pg.get_local_success(work) for work in works): + raise AssertionError(f"injected {operation} unexpectedly succeeded") + if fault.injected_count == 0: + raise AssertionError("transfer fault was not injected") + + active_ranks = pg.get_active_ranks(backend).cpu().tolist() + assert active_ranks == [1, 0], ( + "auto sync should make the transiently failed rank inactive " + f"before {operation} completion, got {active_ranks}" + ) + + +def _inplace_rejoin_worker( + ctx: MooncakePGWorkerContext, + operation: str, + fault_library: str, + fault_started: mp.Event, +) -> None: + """Temporarily fail rank 1's data plane, then rejoin the same process.""" + if ctx.world_size != 2: + raise AssertionError("in-place rejoin test expects exactly two ranks") + + device = ctx.init_group() + backend = ctx.get_backend() + + # 1. Establish a healthy baseline on the selected data path. + _run_rejoin_operation(ctx, device, operation, verify=True) + epoch_before_failure = pg.get_current_epoch(backend) + + # 2. Fail rank 1's data plane. Auto sync must make it inactive. + _run_fault_round( + ctx, device, backend, operation, fault_library, fault_started + ) + + # 3. At an application-selected safe point, the failed rank declares + # itself ready to rejoin. + if ctx.rank == BROKEN_RANK: + pg.join_group(backend) + else: + wait_until( + lambda: pg.get_current_epoch(backend) > epoch_before_failure, + timeout_s=60.0, + poll_interval_s=0.05, + description="waiting for the post-failure group view", + ) + wait_until( + lambda: pg.get_peer_state(backend, [BROKEN_RANK])[0], + timeout_s=30.0, + poll_interval_s=0.05, + description="waiting for the original rank to become activatable", + ) + response = pg.activate_ranks(backend, [BROKEN_RANK]) + assert response.status == pg.ProposalStatus.Applied, ( + "in-place activation should apply, " + f"got {response.status}: {response.reject_reason}" + ) + + # 4. Verify the same communicator can use the same data path again. + _run_rejoin_operation(ctx, device, operation, verify=True) + ctx.record_result({"operation": operation}) + + def _manual_deactivate_recovery_worker( ctx: MooncakePGWorkerContext, broken_exited: mp.Event, @@ -994,6 +1242,61 @@ def test_recovery_after_graceful_group_destroy(self) -> None: self.assertEqual(len(replacement_rows), 1) self.assertEqual(len(removed_rows), 1) + def _run_recovery_p2p(self, graceful_group_destroy: bool) -> None: + recovery_world_size = 2 + spawn_ctx = mp.get_context("spawn") + broken_exited = spawn_ctx.Event() + start_recovery = spawn_ctx.Event() + + rows = self.spawn_backend_and_collect( + _recovery_p2p_worker, + broken_exited, + start_recovery, + graceful_group_destroy, + world_size=recovery_world_size, + nprocs=recovery_world_size + 1, + timeout_s=75.0, + process_exit_events=( + {BROKEN_RANK: broken_exited} + if graceful_group_destroy + else None + ), + ) + + self.assert_all_ok(rows) + + def test_recovery_p2p(self) -> None: + """P2P remains usable after an abruptly failed rank is replaced.""" + self._run_recovery_p2p(graceful_group_destroy=False) + + def test_recovery_p2p_after_graceful_group_destroy(self) -> None: + """P2P remains usable after a gracefully removed rank is replaced.""" + self._run_recovery_p2p(graceful_group_destroy=True) + + def _run_inplace_rejoin(self, operation: str) -> None: + spawn_ctx = mp.get_context("spawn") + + with preload_transfer_fault() as fault_library: + fault_started = spawn_ctx.Event() + rows = self.spawn_backend_and_collect( + _inplace_rejoin_worker, + operation, + str(fault_library), + fault_started, + world_size=2, + nprocs=2, + timeout_s=75.0, + ) + self.assert_all_ok(rows) + + def test_inplace_rejoin_collective(self) -> None: + """Collectives recover when the same live process rejoins.""" + self._run_inplace_rejoin("collective") + + def test_inplace_rejoin_p2p(self) -> None: + """P2P recovers when the same live process rejoins.""" + self._run_inplace_rejoin("p2p") + def test_extension(self) -> None: """Test extension mode allows new ranks to join existing group.""" spawn_ctx = mp.get_context("spawn") diff --git a/mooncake-pg/tests/test_pg_p2p.py b/mooncake-pg/tests/test_pg_p2p.py index 4fac899dd2..f8b151ca87 100644 --- a/mooncake-pg/tests/test_pg_p2p.py +++ b/mooncake-pg/tests/test_pg_p2p.py @@ -151,6 +151,7 @@ def p2p_all_to_all(ctx, device): return dist.batch_isend_irecv(ops) device = ctx.init_group() + backend = ctx.get_backend() # Round 1: all healthy works = p2p_all_to_all(ctx, device) @@ -191,6 +192,14 @@ def p2p_all_to_all(ctx, device): assert not pg.get_local_success(w), \ f"rank {ctx.rank} round 2: P2P with broken peer should fail locally" + expected_active_ranks = [1] * ctx.world_size + expected_active_ranks[BROKEN_RANK] = 0 + active_ranks = pg.get_active_ranks(backend).cpu().tolist() + assert active_ranks == expected_active_ranks, ( + "auto_sync_on_failure should apply the group view before " + f"failed P2P completion, got active_ranks={active_ranks}" + ) + ctx.record_result({"role": "survivor"}) diff --git a/mooncake-pg/tests/transfer_fault_injection.py b/mooncake-pg/tests/transfer_fault_injection.py new file mode 100644 index 0000000000..907c0e30ab --- /dev/null +++ b/mooncake-pg/tests/transfer_fault_injection.py @@ -0,0 +1,231 @@ +"""Test-only transfer fault injection for Mooncake PG workers.""" + +import ctypes +import os +import subprocess +import sys +import tempfile +import unittest +from collections.abc import Generator, Iterable +from contextlib import contextmanager +from pathlib import Path + +from pg_test_utils import temporary_env + +RankPair = tuple[int, int] + + +def _build_preload(output: Path) -> None: + source = Path(__file__).with_name("transfer_fault_preload.cpp") + repository_root = Path(__file__).resolve().parents[2] + transfer_engine_include = ( + repository_root / "mooncake-transfer-engine" / "include" + ) + process_group_include = repository_root / "mooncake-pg" / "include" + command = [ + os.environ.get("CXX", "c++"), + "-std=c++20", + "-O2", + "-shared", + "-fPIC", + f"-I{transfer_engine_include}", + f"-I{process_group_include}", + str(source), + "-ldl", + "-o", + str(output), + ] + try: + subprocess.run(command, check=True, capture_output=True, text=True) + except OSError as error: + raise unittest.SkipTest( + f"transfer fault injection cannot invoke the compiler: {error}" + ) from error + except subprocess.CalledProcessError as error: + reason = (error.stderr or "").strip() + raise unittest.SkipTest( + "transfer fault injection cannot build its preload shim: " + f"{reason or error}" + ) from error + + +def _verify_preload(library: Path, preload: str) -> None: + """Verify LD_PRELOAD and symbol resolution in a fresh process.""" + probe = """ +import ctypes +import sys + +from mooncake import pg + +library = ctypes.CDLL(sys.argv[1]) +available = library.mooncakePgTestFaultAvailable +available.argtypes = [] +available.restype = ctypes.c_int +if not available(): + print("required fault-injection symbols not found", file=sys.stderr) + raise SystemExit(1) +clear_targets = library.mooncakePgTestClearFailedTargets +clear_targets.argtypes = [] +clear_targets.restype = None +add_target = library.mooncakePgTestAddFailedTarget +add_target.argtypes = [ctypes.c_int] +add_target.restype = None +setter = library.mooncakePgTestSetFaultEnabled +setter.argtypes = [ctypes.c_int] +setter.restype = None +reset_counter = library.mooncakePgTestResetFailureCount +reset_counter.argtypes = [] +reset_counter.restype = None +counter = library.mooncakePgTestGetFailureCount +counter.argtypes = [] +counter.restype = ctypes.c_uint64 +setter(0) +clear_targets() +add_target(1) +reset_counter() +setter(1) +setter(0) +clear_targets() +counter() +""" + environment = os.environ.copy() + environment["LD_PRELOAD"] = preload + try: + result = subprocess.run( + [sys.executable, "-c", probe, str(library)], + check=False, + capture_output=True, + text=True, + timeout=30.0, + env=environment, + ) + except (OSError, subprocess.TimeoutExpired) as error: + raise unittest.SkipTest( + f"transfer fault injection is unavailable: {error}" + ) from error + + if result.returncode != 0: + reason = result.stderr.strip() or result.stdout.strip() + if not reason: + reason = f"exit code {result.returncode}" + raise unittest.SkipTest( + f"transfer fault injection is unavailable: {reason}" + ) + + +@contextmanager +def preload_transfer_fault() -> Generator[Path, None, None]: + """Preload the transfer fault shim into spawned workers.""" + if not sys.platform.startswith("linux"): + raise unittest.SkipTest("LD_PRELOAD fault injection requires Linux") + + with tempfile.TemporaryDirectory(prefix="mooncake-pg-fault-") as temp_dir: + library = Path(temp_dir) / "libmooncake_pg_fault.so" + _build_preload(library) + + preload_entries = [str(library)] + if previous_preload := os.environ.get("LD_PRELOAD"): + preload_entries.append(previous_preload) + + preload = os.pathsep.join(preload_entries) + _verify_preload(library, preload) + + with temporary_env({"LD_PRELOAD": preload}): + yield library + + +class TransferFault: + """Coordinate directed link failures across worker processes.""" + + def __init__(self, library: str | Path, *, local_rank: int) -> None: + self._library = ctypes.CDLL(str(library)) + is_available = self._library.mooncakePgTestFaultAvailable + is_available.argtypes = [] + is_available.restype = ctypes.c_int + if not is_available(): + raise RuntimeError( + "preloaded transfer fault shim cannot resolve its " + "required symbols" + ) + self._clear_failed_targets = ( + self._library.mooncakePgTestClearFailedTargets + ) + self._clear_failed_targets.argtypes = [] + self._clear_failed_targets.restype = None + self._add_failed_target = self._library.mooncakePgTestAddFailedTarget + self._add_failed_target.argtypes = [ctypes.c_int] + self._add_failed_target.restype = None + self._set_enabled = self._library.mooncakePgTestSetFaultEnabled + self._set_enabled.argtypes = [ctypes.c_int] + self._set_enabled.restype = None + self._reset_count = self._library.mooncakePgTestResetFailureCount + self._reset_count.argtypes = [] + self._reset_count.restype = None + self._get_count = self._library.mooncakePgTestGetFailureCount + self._get_count.argtypes = [] + self._get_count.restype = ctypes.c_uint64 + if local_rank < 0: + raise ValueError("local rank must be non-negative") + self._local_rank = local_rank + self._set_enabled(0) + self._clear_failed_targets() + self._reset_count() + self._active = False + + @staticmethod + def _normalize_links( + links: Iterable[RankPair], + ) -> tuple[RankPair, ...]: + normalized = tuple(dict.fromkeys(links)) + if not normalized: + raise ValueError("at least one failed link is required") + for source_rank, target_rank in normalized: + if ( + source_rank < 0 + or target_rank < 0 + or source_rank == target_rank + ): + raise ValueError( + "failed links require distinct, non-negative ranks" + ) + return normalized + + @contextmanager + def failing_links( + self, links: Iterable[RankPair] + ) -> Generator[None, None, None]: + """Fail directed links whose source is this worker's rank.""" + if self._active: + raise RuntimeError("fault scopes cannot be nested") + normalized = self._normalize_links(links) + local_targets = { + target_rank + for source_rank, target_rank in normalized + if source_rank == self._local_rank + } + + self._set_enabled(0) + self._clear_failed_targets() + for target_rank in local_targets: + self._add_failed_target(target_rank) + self._reset_count() + self._set_enabled(bool(local_targets)) + self._active = True + try: + yield + finally: + self._set_enabled(0) + self._clear_failed_targets() + self._active = False + + @contextmanager + def failing_link( + self, source_rank: int, target_rank: int + ) -> Generator[None, None, None]: + """Fail one directed global-rank link within this scope.""" + with self.failing_links(((source_rank, target_rank),)): + yield + + @property + def injected_count(self) -> int: + return int(self._get_count()) diff --git a/mooncake-pg/tests/transfer_fault_preload.cpp b/mooncake-pg/tests/transfer_fault_preload.cpp new file mode 100644 index 0000000000..fe8ae8195e --- /dev/null +++ b/mooncake-pg/tests/transfer_fault_preload.cpp @@ -0,0 +1,189 @@ +// Copyright 2026 KVCache.AI +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "control_plane/link_manager.h" +#include "transfer_engine.h" + +namespace { + +std::atomic fault_enabled{false}; +std::atomic injected_failure_count{0}; + +std::mutex fault_config_mutex; +std::set failed_targets; + +// Fault rules use global ranks, while transfer requests identify targets by +// segment ID. Capture both mappings so we can apply rank-level rules to +// individual transfer tasks. +std::mutex mapping_mutex; +std::unordered_map + segment_id_to_rank; +std::unordered_map> + batch_id_to_target_ranks; + +constexpr char kGetTransferStatusSymbol[] = + "_ZN8mooncake14TransferEngine17getTransferStatusEmmRNS_" + "9Transport14TransferStatusE"; +constexpr char kSubmitTransferSymbol[] = + "_ZN8mooncake14TransferEngine14submitTransferEmRKSt6vectorINS_" + "9Transport15TransferRequestESaIS3_EE"; +constexpr char kResolvePeerSymbol[] = + "_ZNK8mooncake11LinkManager11resolvePeerEi"; + +using GetTransferStatus = mooncake::Status (*)(mooncake::TransferEngine*, + mooncake::BatchID, size_t, + mooncake::TransferStatus&); +using SubmitTransfer = + mooncake::Status (*)(mooncake::TransferEngine*, mooncake::BatchID, + const std::vector&); +using ResolvePeer = std::optional (*)( + const mooncake::LinkManager*, mooncake::GlobalRank); + +void* findRealSymbol(const char* name) { + if (auto symbol = dlsym(RTLD_NEXT, name)) return symbol; + + // Python loads extension dependencies in a local ELF scope, where + // RTLD_NEXT cannot see libmooncake_pg even though interposition still + // applies to its PLT calls. Look up the original definition directly + // from the already-loaded library in that case. + static auto handle = dlopen("libmooncake_pg.so", RTLD_LAZY | RTLD_NOLOAD); + return handle ? dlsym(handle, name) : nullptr; +} + +GetTransferStatus findRealGetTransferStatus() { + static auto real = reinterpret_cast( + findRealSymbol(kGetTransferStatusSymbol)); + return real; +} + +SubmitTransfer findRealSubmitTransfer() { + static auto real = + reinterpret_cast(findRealSymbol(kSubmitTransferSymbol)); + return real; +} + +ResolvePeer findRealResolvePeer() { + static auto real = + reinterpret_cast(findRealSymbol(kResolvePeerSymbol)); + return real; +} + +bool shouldInjectFailure(mooncake::GlobalRank target_rank) { + if (!fault_enabled.load(std::memory_order_acquire)) return false; + + std::lock_guard lock(fault_config_mutex); + return failed_targets.contains(target_rank); +} + +} // namespace + +namespace mooncake { + +std::optional LinkManager::resolvePeer( + GlobalRank peer) const { + auto real = findRealResolvePeer(); + if (!real) std::abort(); + auto target_id = real(this, peer); + if (target_id.has_value()) { + std::lock_guard lock(mapping_mutex); + segment_id_to_rank[*target_id] = peer; + } + return target_id; +} + +Status TransferEngine::submitTransfer( + BatchID batch_id, const std::vector& entries) { + auto real = findRealSubmitTransfer(); + if (!real) std::abort(); + auto result = real(this, batch_id, entries); + if (result.ok()) { + std::vector target_ranks(entries.size(), + kInvalidGlobalRank); + std::lock_guard lock(mapping_mutex); + for (size_t task_id = 0; task_id < entries.size(); ++task_id) { + auto rank_it = segment_id_to_rank.find(entries[task_id].target_id); + if (rank_it != segment_id_to_rank.end()) { + target_ranks[task_id] = rank_it->second; + } + } + batch_id_to_target_ranks.insert_or_assign(batch_id, + std::move(target_ranks)); + } + return result; +} + +Status TransferEngine::getTransferStatus(BatchID batch_id, size_t task_id, + TransferStatus& status) { + auto real = findRealGetTransferStatus(); + if (!real) std::abort(); + auto result = real(this, batch_id, task_id, status); + if (result.ok() && status.s == TransferStatusEnum::COMPLETED && + fault_enabled.load(std::memory_order_acquire)) { + GlobalRank target_rank = kInvalidGlobalRank; + { + std::lock_guard lock(mapping_mutex); + auto batch_it = batch_id_to_target_ranks.find(batch_id); + if (batch_it != batch_id_to_target_ranks.end() && + task_id < batch_it->second.size()) { + target_rank = batch_it->second[task_id]; + } + } + if (shouldInjectFailure(target_rank)) { + status.s = TransferStatusEnum::FAILED; + injected_failure_count.fetch_add(1, std::memory_order_relaxed); + } + } + return result; +} + +} // namespace mooncake + +extern "C" int mooncakePgTestFaultAvailable() { + return findRealGetTransferStatus() != nullptr && + findRealSubmitTransfer() != nullptr && + findRealResolvePeer() != nullptr; +} + +extern "C" void mooncakePgTestClearFailedTargets() { + std::lock_guard lock(fault_config_mutex); + failed_targets.clear(); +} + +extern "C" void mooncakePgTestAddFailedTarget(int target_rank) { + std::lock_guard lock(fault_config_mutex); + failed_targets.insert(target_rank); +} + +extern "C" void mooncakePgTestSetFaultEnabled(int enabled) { + fault_enabled.store(enabled != 0, std::memory_order_release); +} + +extern "C" void mooncakePgTestResetFailureCount() { + injected_failure_count.store(0, std::memory_order_release); +} + +extern "C" uint64_t mooncakePgTestGetFailureCount() { + return injected_failure_count.load(std::memory_order_acquire); +} From 86b21ccfe5f528e8c98f344737ee2eaf18ae7b20 Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Sat, 8 Aug 2026 22:19:00 +0800 Subject: [PATCH 005/483] [Wheel] Optimize typed-ragged rollout layout (#3113) Restore the PR2671 flat-dict GET fast path for encoded non-tensor rollout data. Flat dict reads now avoid the DataProto object-array round trip, msgpack ragged values use streaming decode again, and encoded payload members are materialized by stage. Add regression coverage for dict-vs-DataProto output shape and direct-copy typed-ragged payload tests. --- .../mooncake/structured_object_store.py | 639 ++++++++++++++---- .../tests/test_structured_object_store.py | 232 ++++++- 2 files changed, 741 insertions(+), 130 deletions(-) diff --git a/mooncake-wheel/mooncake/structured_object_store.py b/mooncake-wheel/mooncake/structured_object_store.py index d2a13e5223..6eee7df21e 100644 --- a/mooncake-wheel/mooncake/structured_object_store.py +++ b/mooncake-wheel/mooncake/structured_object_store.py @@ -31,6 +31,8 @@ -704 ) # Mooncake remove returns -704 for an already-missing object. STRUCTURED_FIELD_SPECS_KEY = "__mooncake_structured_fields__" +_TYPED_RAGGED_DEFAULT_DTYPE = "int64" +_RAGGED_TENSOR_DEFAULT_DTYPE = "torch.float32" _ENCODING_FALLBACK_ERRORS = ( TypeError, ValueError, @@ -666,7 +668,7 @@ def get( """Materialize a DataProto-like object or flat dict.""" if type not in {"dataproto", "dict"}: raise ValueError(f"unsupported Mooncake payload type: {type!r}") - result = self.get_dataproto( + result = self._get_dataproto( ref, fields=fields, batch_fields=batch_fields, @@ -675,6 +677,7 @@ def get( data_cls=dict if type == "dict" else data_cls, destinations=destinations, rows=rows, + _flat_dict_output=(type == "dict"), ) return _envelope_to_flat_dict(result) if type == "dict" else result @@ -746,6 +749,31 @@ def get_dataproto( data_cls: Optional[Any] = None, destinations: Optional[Mapping[str, Any]] = None, rows: slice | StructuredMemberSlice | Sequence[int] | None = None, + ) -> Any: + """Materialize selected DataProto fields from structured object refs.""" + return self._get_dataproto( + ref, + fields=fields, + batch_fields=batch_fields, + non_tensor_fields=non_tensor_fields, + meta_info_keys=meta_info_keys, + data_cls=data_cls, + destinations=destinations, + rows=rows, + ) + + def _get_dataproto( + self, + ref: DataProtoRefLike, + *, + fields: Optional[Sequence[str]] = None, + batch_fields: Optional[Sequence[str]] = None, + non_tensor_fields: Optional[Sequence[str]] = None, + meta_info_keys: Optional[Sequence[str]] = None, + data_cls: Optional[Any] = None, + destinations: Optional[Mapping[str, Any]] = None, + rows: slice | StructuredMemberSlice | Sequence[int] | None = None, + _flat_dict_output: bool = False, ) -> Any: """Materialize selected DataProto fields from structured object refs.""" ref = _resolve_dataproto_ref(ref) @@ -812,41 +840,61 @@ def get_dataproto( batch[name] = value else: non_tensor_batch[name] = value - for name, location in encoded_requests: - stage_ref = ref.stage_refs[location.stage] - encoded = ref.encoded_non_tensor[name] - if row_selection is None: - members = list(encoded["payload_members"].values()) + if row_selection is None and encoded_requests: + by_stage_encoded: dict[str, list[tuple[str, StructuredFieldLocation]]] = {} + for name, location in encoded_requests: + by_stage_encoded.setdefault(location.stage, []).append((name, location)) + for stage, entries in by_stage_encoded.items(): + stage_ref = ref.stage_refs[stage] + members: list[str] = [] + for name, _location in entries: + encoded = ref.encoded_non_tensor[name] + members.extend(encoded["payload_members"].values()) result = self.materialize( self.read_spec(stage_ref).select_members(members) ) - payload = { - payload_name: result.objects[member] - for payload_name, member in encoded["payload_members"].items() - } - values = _decode_structured_non_tensor_encoded( - encoded, payload, ref.batch_size, encoded.get("metadata") - ) - elif row_indices is not None: - payload, metadata = self._read_structured_non_tensor_payload_indices( - stage_ref, - encoded, - row_indices, - ) - values = _decode_structured_non_tensor_encoded( - encoded, payload, output_rows, metadata - ) - else: - payload, metadata = self._read_structured_non_tensor_payload_slice( - stage_ref, - encoded, - row_slice, - ref.batch_size, - ) - values = _decode_structured_non_tensor_encoded( - encoded, payload, output_rows, metadata + for name, _location in entries: + encoded = ref.encoded_non_tensor[name] + payload = { + payload_name: result.objects[member] + for payload_name, member in encoded["payload_members"].items() + } + values = _decode_structured_non_tensor_encoded( + encoded, payload, ref.batch_size, encoded.get("metadata") + ) + non_tensor_batch[name] = ( + values + if _flat_dict_output + else _object_array_from_decoded_values(values) + ) + else: + for name, location in encoded_requests: + stage_ref = ref.stage_refs[location.stage] + encoded = ref.encoded_non_tensor[name] + if row_indices is not None: + payload, metadata = self._read_structured_non_tensor_payload_indices( + stage_ref, + encoded, + row_indices, + ) + values = _decode_structured_non_tensor_encoded( + encoded, payload, output_rows, metadata + ) + else: + payload, metadata = self._read_structured_non_tensor_payload_slice( + stage_ref, + encoded, + row_slice, + ref.batch_size, + ) + values = _decode_structured_non_tensor_encoded( + encoded, payload, output_rows, metadata + ) + non_tensor_batch[name] = ( + values + if _flat_dict_output + else _object_array_from_decoded_values(values) ) - non_tensor_batch[name] = _object_array_from_decoded_values(values) meta_info = _select_mapping(ref.meta_info, meta_info_keys) return _build_dataproto_like_result( batch, non_tensor_batch, meta_info, data_cls @@ -1881,7 +1929,7 @@ def _envelope_to_flat_dict(data: Mapping[str, Any]) -> dict[str, Any]: result.update(batch) for name, value in non_tensor_batch.items(): result[name] = ( - list(value) + value.tolist() if isinstance(value, np.ndarray) and value.dtype == object else value ) @@ -2598,81 +2646,57 @@ def _encode_ragged_tensor_values( ) -> tuple[dict[str, Any], dict[str, Any]]: if _torch is None: raise RuntimeError("torch is required to encode ragged tensor fields") - tensors: list[Any] = [] - dtype = None - max_ndim = 0 + # Determine torch dtype and convert all values to numpy arrays. + torch_dtype = None + converted: list[Any] = [] for value in values: if value is None: - tensors.append(None) + converted.append(None) continue - tensor = value.detach() + if isinstance(value, np.ndarray): + if value.dtype == object: + raise ValueError("ragged tensor codec requires numeric ndarray values") + if not value.flags.writeable: + value = value.copy() + tensor = _torch.as_tensor(value) + else: + tensor = value.detach() if tensor.device.type != "cpu" or not tensor.is_contiguous(): tensor = tensor.cpu().contiguous() - dtype = tensor.dtype if dtype is None else dtype - if tensor.dtype != dtype: - raise ValueError(f"mixed tensor dtype: {dtype} vs {tensor.dtype}") - max_ndim = max(max_ndim, tensor.dim()) - tensors.append(tensor) - offsets = _torch.zeros(len(tensors) + 1, dtype=_torch.int64) - ndims = _torch.zeros(len(tensors), dtype=_torch.int16) - shapes = _torch.zeros((len(tensors), max(max_ndim, 1)), dtype=_torch.int64) - nulls = np.asarray([tensor is None for tensor in tensors], dtype=np.bool_) - flat_parts = [] - offset = 0 - for row, tensor in enumerate(tensors): - if tensor is None: - offsets[row + 1] = offset - continue - flat = tensor.reshape(-1) - flat_parts.append(flat) - offset += flat.numel() - offsets[row + 1] = offset - ndims[row] = tensor.dim() - if tensor.dim() > 0: - shapes[row, : tensor.dim()] = _torch.tensor( - list(tensor.shape), dtype=_torch.int64 - ) - data_dtype = dtype or _torch.float32 - data = ( - _torch.cat(flat_parts) if flat_parts else _torch.empty((0,), dtype=data_dtype) - ) - payload = { - "data": data.numpy(), - "offsets": offsets.numpy(), - "shapes": shapes.numpy(), - "ndims": ndims.numpy(), - "nulls": nulls, - } - return payload, { - "dtype": str(data_dtype), - "max_ndim": int(max_ndim), - "shape_policy": "ragged", - } - + if tensor.dtype == _torch.bfloat16: + raise ValueError("ragged tensor codec does not support torch.bfloat16") + torch_dtype = tensor.dtype if torch_dtype is None else torch_dtype + if tensor.dtype != torch_dtype: + raise ValueError(f"mixed tensor dtype: {torch_dtype} vs {tensor.dtype}") + converted.append(tensor.numpy()) + data_dtype = torch_dtype or _parse_torch_dtype(_RAGGED_TENSOR_DEFAULT_DTYPE) + np_dtype = np.dtype(_torch_dtype_to_numpy(str(data_dtype))) + payload, meta = _encode_typed_ragged_values(converted, np_dtype) + meta["dtype"] = str(data_dtype) + return payload, meta def _decode_ragged_tensor_values( - payload: dict[str, Any], - rows: int, - metadata: Optional[Mapping[str, Any]] = None, + payload: dict[str, Any], rows: int, metadata: Optional[Mapping[str, Any]] = None ) -> list[Any]: if _torch is None: raise RuntimeError("torch is required to decode ragged tensor fields") - data = _torch.from_numpy(payload["data"]) - offsets = payload["offsets"] - shapes = payload["shapes"] - ndims = payload["ndims"] - nulls = payload["nulls"] - values = [] - for row in range(rows): - if bool(nulls[row]): - values.append(None) - continue - begin = int(offsets[row]) - end = int(offsets[row + 1]) - ndim = int(ndims[row]) - shape = tuple(int(v) for v in shapes[row, :ndim].tolist()) - values.append(data[begin:end].reshape(shape)) - return values + # Metadata stores torch dtype strings (e.g. "torch.int64"); convert to numpy + # dtype so _decode_typed_ragged_values can parse it. + dtype_str = (metadata or {}).get("dtype", _RAGGED_TENSOR_DEFAULT_DTYPE) + np_dtype = _torch_dtype_to_numpy(dtype_str) + patched_meta = dict(metadata) if metadata else {} + patched_meta["dtype"] = str(np_dtype) + np_values = _decode_typed_ragged_values(payload, rows, patched_meta) + torch_dtype = _parse_torch_dtype(dtype_str) + result: list[Any] = [] + for v in np_values: + if v is None: + result.append(None) + elif isinstance(v, np.ndarray): + result.append(_torch.from_numpy(v).to(torch_dtype)) + else: + result.append(_torch.as_tensor(v, dtype=torch_dtype)) + return result def _normalize_ragged_tensor_dict_keys(keys: Any) -> list[str]: if isinstance(keys, Mapping): @@ -2819,15 +2843,20 @@ def _encode_typed_ragged_values( ) -> tuple[dict[str, Any], dict[str, Any]]: if dtype_hint is None: source_arrays = [np.asarray(value) for value in values if value is not None] - dtype = np.result_type(*source_arrays) if source_arrays else np.dtype(np.int64) + dtype = np.result_type(*source_arrays) if source_arrays else np.dtype(_TYPED_RAGGED_DEFAULT_DTYPE) else: dtype = np.dtype(dtype_hint) if dtype.hasobject: raise ValueError("typed_ragged codec requires non-object dtype") + + ndarray_encoded = _encode_typed_ragged_ndarray_rows(values, dtype) + if ndarray_encoded is not None: + return ndarray_encoded + arrays = [ np.asarray([], dtype=dtype) if value is None - else np.ascontiguousarray(np.asarray(value, dtype=dtype)) + else _as_contiguous_array_preserve_ndim(value, dtype) for value in values ] max_ndim = max((array.ndim for array in arrays), default=0) @@ -2845,27 +2874,33 @@ def _encode_typed_ragged_values( ndims[row] = array.ndim if array.ndim > 0: shapes[row, : array.ndim] = array.shape + total_elems = int(offset) if flat_arrays: if _concat_arrays_into is not None: - data = _DirectCopyPayload.from_flat_arrays( - flat_arrays, dtype, int(offset) - ) + data = _DirectCopyPayload.from_flat_arrays(flat_arrays, dtype, total_elems) else: buffers = tuple(memoryview(flat.data).cast("B") for flat in flat_arrays) data = _MultiBufferPayload( buffers=buffers, owners=tuple(flat_arrays), dtype=np.dtype(dtype).str, - shape=(int(offset),), + shape=(total_elems,), ) else: empty = np.empty(0, dtype=dtype) data = _MultiBufferPayload( - buffers=(memoryview(empty.data).cast("B"),), - owners=(empty,), - dtype=np.dtype(dtype).str, - shape=(0,), + buffers=(memoryview(empty.data).cast("B"),), owners=(empty,), + dtype=np.dtype(dtype).str, shape=(0,), ) + ndarray_rows = [ + value is not None and isinstance(value, np.ndarray) for value in values + ] + if all(value is None or is_array for value, is_array in zip(values, ndarray_rows)): + row_format = "ndarray" + elif any(ndarray_rows): + row_format = "mixed" + else: + row_format = "list" return ( { "data": data, @@ -2874,18 +2909,47 @@ def _encode_typed_ragged_values( "ndims": ndims, "nulls": nulls, }, - {"dtype": str(dtype), "max_ndim": int(max_ndim), "shape_policy": "ragged"}, + { + "dtype": str(dtype), + "max_ndim": int(max_ndim), + "shape_policy": "ragged", + "row_format": row_format, + "ndarray_rows": ndarray_rows if row_format == "mixed" else None, + }, ) - def _decode_typed_ragged_values( payload: dict[str, Any], rows: int, metadata: Optional[Mapping[str, Any]] = None ) -> list[Any]: - data = payload["data"] + raw = payload["data"] + dtype_str = (metadata or {}).get("dtype", _TYPED_RAGGED_DEFAULT_DTYPE) + np_dtype = np.dtype(dtype_str) + # data may arrive as torch.Tensor, bytes, or numpy (pool-backed or typed) + if _torch is not None and isinstance(raw, _torch.Tensor): + data = raw.numpy() + elif isinstance(raw, (bytes, bytearray, memoryview)): + data = np.frombuffer(raw, dtype=np_dtype) + elif isinstance(raw, np.ndarray): + data = raw if raw.dtype == np_dtype else np.frombuffer(raw, dtype=np_dtype) + else: + data = np.asarray(raw, dtype=np_dtype) offsets = payload["offsets"] shapes = payload["shapes"] ndims = payload["ndims"] nulls = payload["nulls"] + metadata = metadata or {} + row_format = metadata.get("row_format") + if ( + row_format == "ndarray" + and metadata.get("physical_layout") == "contiguous_flat" + ): + fast_values = _decode_typed_ragged_ndarray_rows_fast( + data, offsets, shapes, ndims, nulls, rows, metadata + ) + if fast_values is not None: + return fast_values + + ndarray_rows = metadata.get("ndarray_rows") or [] values = [] for row in range(rows): if bool(nulls[row]): @@ -2895,7 +2959,13 @@ def _decode_typed_ragged_values( end = int(offsets[row + 1]) ndim = int(ndims[row]) shape = tuple(int(v) for v in shapes[row, :ndim].tolist()) - values.append(data[begin:end].reshape(shape).tolist()) + array = data[begin:end].reshape(shape) + if row_format == "ndarray" or ( + row_format == "mixed" and bool(ndarray_rows[row]) + ): + values.append(array) + else: + values.append(array.tolist()) return values @@ -3695,7 +3765,6 @@ def _put_tensor_payload( memoryview(payload), len(payload) or 1, transfer_policy, - pre_registered=False, config=config, ) payload_spec["metadata_bytes"] = metadata_bytes @@ -5338,7 +5407,313 @@ def _normalize_structured_scalar(value: Any) -> Any: +def _encode_typed_ragged_regular_ndarray_rows( + values: list[Any], dtype: np.dtype[Any] +) -> tuple[dict[str, Any], dict[str, Any]] | None: + row_count = len(values) + if row_count < 2: + return None + first = values[0] + if ( + not isinstance(first, np.ndarray) + or first.dtype != dtype + or not first.flags.c_contiguous + ): + return None + shape = first.shape + for value in values[1:]: + if ( + not isinstance(value, np.ndarray) + or value.dtype != dtype + or not value.flags.c_contiguous + or value.shape != shape + ): + return None + + row_elems = int(first.size) + total_elems = row_count * row_elems + if row_elems: + flat_arrays = [value.reshape(-1) for value in values] + if _concat_arrays_into is not None: + data = _DirectCopyPayload.from_flat_arrays(flat_arrays, dtype, total_elems) + else: + data = _MultiBufferPayload( + buffers=tuple(memoryview(flat.data).cast("B") for flat in flat_arrays), + owners=tuple(flat_arrays), + dtype=np.dtype(dtype).str, + shape=(total_elems,), + ) + else: + data = np.asarray(values, dtype=dtype).reshape(-1) + offsets = np.arange(row_count + 1, dtype=np.int64) * row_elems + nulls = np.zeros(row_count, dtype=np.bool_) + max_ndim = int(first.ndim) + if max_ndim == 0: + shapes = np.zeros((row_count, 1), dtype=np.int64) + ndims = np.zeros(row_count, dtype=np.int16) + else: + shapes = np.empty((row_count, max_ndim), dtype=np.int64) + shapes[:] = shape + ndims = np.full(row_count, max_ndim, dtype=np.int16) + return ( + { + "data": data, + "offsets": offsets, + "shapes": shapes, + "ndims": ndims, + "nulls": nulls, + }, + { + "dtype": str(dtype), + "max_ndim": max_ndim, + "shape_policy": "ragged", + "row_format": "ndarray", + "ndarray_rows": None, + "physical_layout": "contiguous_flat", + }, + ) + + +def _encode_typed_ragged_ndarray_rows( + values: list[Any], dtype: np.dtype[Any] +) -> tuple[dict[str, Any], dict[str, Any]] | None: + regular_encoded = _encode_typed_ragged_regular_ndarray_rows(values, dtype) + if regular_encoded is not None: + return regular_encoded + + row_count = len(values) + if row_count == 0: + return None + + nulls = np.empty(row_count, dtype=np.bool_) + lengths = np.empty(row_count, dtype=np.int64) + arrays: list[np.ndarray] = [] + tail_shape: tuple[int, ...] | None = None + row_ndim: int | None = None + tail_compatible = True + saw_scalar = False + saw_non_scalar = False + + for row, value in enumerate(values): + if value is None: + nulls[row] = True + lengths[row] = 0 + continue + if not isinstance(value, np.ndarray): + return None + + array = _as_contiguous_array_preserve_ndim(value, dtype) + arrays.append(array) + nulls[row] = False + lengths[row] = array.size + if array.ndim == 0: + saw_scalar = True + if saw_non_scalar: + tail_compatible = False + continue + + saw_non_scalar = True + if saw_scalar: + tail_compatible = False + current_tail = array.shape[1:] + if tail_shape is None: + tail_shape = current_tail + row_ndim = array.ndim + elif array.ndim != row_ndim or current_tail != tail_shape: + tail_compatible = False + + if tail_compatible and tail_shape is not None and any(dim == 0 for dim in tail_shape): + tail_compatible = False + + offsets = np.empty(row_count + 1, dtype=np.int64) + offsets[0] = 0 + np.cumsum(lengths, out=offsets[1:]) + total_elems = int(offsets[-1]) + if len(arrays) == 1: + flat = arrays[0].reshape(-1) + data = _MultiBufferPayload( + buffers=(memoryview(flat.data).cast("B"),), + owners=(flat,), + dtype=np.dtype(dtype).str, + shape=(total_elems,), + ) + elif total_elems: + flat_arrays = [arr.reshape(-1) for arr in arrays] + if _concat_arrays_into is not None: + data = _DirectCopyPayload.from_flat_arrays(flat_arrays, dtype, total_elems) + else: + buffers = tuple(memoryview(flat.data).cast("B") for flat in flat_arrays) + data = _MultiBufferPayload( + buffers=buffers, + owners=tuple(flat_arrays), + dtype=np.dtype(dtype).str, + shape=(total_elems,), + ) + else: + empty = np.empty(0, dtype=dtype) + data = _MultiBufferPayload( + buffers=(memoryview(empty.data).cast("B"),), + owners=(empty,), + dtype=np.dtype(dtype).str, + shape=(0,), + ) + + if not arrays or saw_scalar and not saw_non_scalar: + max_ndim = 0 + elif tail_compatible: + max_ndim = int(row_ndim or 1) + else: + max_ndim = max(array.ndim for array in arrays) + + ndims = np.zeros(row_count, dtype=np.int16) + if max_ndim == 0: + shapes = np.zeros((row_count, 1), dtype=np.int64) + elif tail_compatible: + shapes = np.zeros((row_count, max_ndim), dtype=np.int64) + tail = tail_shape or () + tail_elems = int(np.prod(tail, dtype=np.int64)) if tail else 1 + shapes[:, 0] = lengths // tail_elems + if max_ndim > 1: + shapes[:, 1:] = tail + ndims[~nulls] = max_ndim + else: + shapes = np.zeros((row_count, max(max_ndim, 1)), dtype=np.int64) + array_index = 0 + for row, is_null in enumerate(nulls): + if bool(is_null): + continue + array = arrays[array_index] + array_index += 1 + ndims[row] = array.ndim + if array.ndim > 0: + shapes[row, : array.ndim] = array.shape + + return ( + { + "data": data, + "offsets": offsets, + "shapes": shapes, + "ndims": ndims, + "nulls": nulls, + }, + { + "dtype": str(dtype), + "max_ndim": int(max_ndim), + "shape_policy": "ragged", + "row_format": "ndarray", + "ndarray_rows": None, + "physical_layout": "contiguous_flat", + }, + ) + + +def _as_contiguous_array_preserve_ndim(value: Any, dtype: np.dtype[Any]) -> np.ndarray: + if isinstance(value, np.ndarray) and value.dtype == dtype and value.flags.c_contiguous: + return value + array = np.asarray(value, dtype=dtype) + if array.flags.c_contiguous: + return array + return np.array(array, dtype=dtype, order="C", copy=True) + + +def _decode_typed_ragged_ndarray_rows_fast( + data: np.ndarray, + offsets: np.ndarray, + shapes: np.ndarray, + ndims: np.ndarray, + nulls: np.ndarray, + rows: int, + metadata: Mapping[str, Any], +) -> list[Any] | None: + # Returned ndarray rows are views into data; _OwnerBackedList preserves any + # pool owner attached to that data for the lifetime of the row list. + max_ndim = int(metadata.get("max_ndim", 0)) + owner = getattr(data, "_mooncake_pool_owner", None) + source = data.view(np.ndarray) if isinstance(data, np.ndarray) else data + has_nulls = bool(nulls.any()) + + def attach_owner(values: list[Any]) -> list[Any]: + if owner is None: + return values + return _OwnerBackedList(values, owner) + + if max_ndim == 1: + if not has_nulls: + row_width = _regular_offsets_width(offsets, rows) + if row_width is not None: + return attach_owner(list(source.reshape((rows, row_width)))) + begins = offsets[:-1].tolist() + ends = offsets[1:].tolist() + return attach_owner( + [source[b:e] for b, e in zip(begins, ends)] + ) + values: list[Any] = [] + for row, is_null in enumerate(nulls): + if bool(is_null): + values.append(None) + else: + begin = int(offsets[row]) + end = int(offsets[row + 1]) + values.append(source[begin:end]) + return attach_owner(values) + if max_ndim <= 1 or rows == 0: + return None + + non_null_rows = np.flatnonzero(~nulls) + if non_null_rows.size == 0: + return [None] * rows + first_row = int(non_null_rows[0]) + if int(ndims[first_row]) != max_ndim: + return None + tail = tuple(int(v) for v in shapes[first_row, 1:max_ndim]) + if any(dim == 0 for dim in tail): + return None + for row in non_null_rows: + row_index = int(row) + if int(ndims[row_index]) != max_ndim: + return None + if tuple(int(v) for v in shapes[row_index, 1:max_ndim]) != tail: + return None + + tail_elems = int(np.prod(tail, dtype=np.int64)) if tail else 1 + if tail_elems <= 0: + return None + if int(offsets[-1]) % tail_elems != 0: + return None + flat_rows = source.reshape((-1, *tail)) + first_axis_offsets = offsets // tail_elems + + if not has_nulls: + row_width = _regular_offsets_width(first_axis_offsets, rows) + if row_width is not None: + return attach_owner(list(flat_rows.reshape((rows, row_width, *tail)))) + fa_begins = first_axis_offsets[:-1].tolist() + fa_ends = first_axis_offsets[1:].tolist() + return attach_owner( + [flat_rows[b:e] for b, e in zip(fa_begins, fa_ends)] + ) + values = [] + for row, is_null in enumerate(nulls): + if bool(is_null): + values.append(None) + else: + begin = int(first_axis_offsets[row]) + end = int(first_axis_offsets[row + 1]) + values.append(flat_rows[begin:end]) + return attach_owner(values) + + +def _regular_offsets_width(offsets: np.ndarray, rows: int) -> int | None: + if rows <= 0 or int(offsets[0]) != 0: + return None + total = int(offsets[-1]) + if total < 0 or total % rows != 0: + return None + row_width = total // rows + if not bool(np.all(offsets[1:] - offsets[:-1] == row_width)): + return None + return row_width def _encode_msgpack_ragged_values( path: str, values: list[Any] @@ -5368,8 +5743,6 @@ def _encode_msgpack_ragged_values( {}, ) - - def _decode_msgpack_ragged_values(payload: dict[str, Any], rows: int) -> list[Any]: data = payload["data"] offsets = payload["offsets"] @@ -5383,14 +5756,24 @@ def _decode_msgpack_ragged_values(payload: dict[str, Any], rows: int) -> list[An f"msgpack_ragged nulls length {len(nulls)} does not match rows {rows}" ) raw_data = bytes(data) if not isinstance(data, bytes) else data + if rows == 0 or not bool(nulls.any()): + unpacker = _msgpack.Unpacker(raw=False) + unpacker.feed(raw_data) + values = list(unpacker) + if len(values) != rows: + raise ValueError( + f"msgpack_ragged decoded {len(values)} rows, expected {rows}" + ) + return values + unpacker = _msgpack.Unpacker(raw=False) + unpacker.feed(raw_data) + non_null_iter = iter(unpacker) values = [] - for row, is_null in enumerate(nulls): + for is_null in nulls: if bool(is_null): values.append(None) else: - begin = int(offsets[row]) - end = int(offsets[row + 1]) - values.append(_msgpack.unpackb(raw_data[begin:end], raw=False)) + values.append(next(non_null_iter)) return values @@ -5399,8 +5782,6 @@ def __init__(self, values: Sequence[Any], owner: Any) -> None: super().__init__(values) self._mooncake_pool_owner = owner - - class _OwnerBackedObjectArray(np.ndarray): def __array_finalize__(self, obj: Any) -> None: if obj is not None: @@ -5413,8 +5794,6 @@ def tolist(self) -> list[Any]: return values return _OwnerBackedList(values, owner) - - def _object_array_from_decoded_values(values: list[Any]) -> np.ndarray: array = np.empty(len(values), dtype=object) array[:] = values @@ -5722,7 +6101,21 @@ def _cleanup_keys(store: BundleStore, keys: Sequence[str], strict: bool) -> None _torch = None # type: ignore[assignment] +def _parse_torch_dtype(dtype_str: str) -> Any: + """Parse torch dtype string (e.g. 'torch.float32') to torch.dtype object.""" + name = dtype_str.removeprefix("torch.") + dtype = getattr(_torch, name, None) + if dtype is None or not isinstance(dtype, _torch.dtype): + raise ValueError(f"unknown torch dtype: {dtype_str}") + return dtype + +def _torch_dtype_to_numpy(dtype_str: str) -> np.dtype: + """Convert value-preserving torch dtypes to numpy dtypes.""" + torch_dtype = _parse_torch_dtype(dtype_str) + if torch_dtype == _torch.bfloat16: + raise ValueError("torch.bfloat16 has no value-preserving numpy dtype") + return _torch.empty(0, dtype=torch_dtype).numpy().dtype @dataclass class _CodecDecision: diff --git a/mooncake-wheel/tests/test_structured_object_store.py b/mooncake-wheel/tests/test_structured_object_store.py index 13131d290d..971c4d89b5 100644 --- a/mooncake-wheel/tests/test_structured_object_store.py +++ b/mooncake-wheel/tests/test_structured_object_store.py @@ -434,6 +434,25 @@ def write_manifest( ) +def test_default_bundle_chunk_tuning_matches_rollout_transfer_target() -> None: + assert sos.DEFAULT_BUNDLE_CHUNK_BYTES == 64 * 1024**2 + assert sos.AUTO_PARALLEL_MIN_BYTES == sos.DEFAULT_BUNDLE_CHUNK_BYTES + + +def test_small_dict_sized_groups_enable_auto_parallel_put() -> None: + class SizedBuffer: + def __init__(self, size: int) -> None: + self._size = size + + def __len__(self) -> int: + return self._size + + _store, transfer = make_transfer() + groups = [[SizedBuffer(1024**2)] for _ in range(128)] + policy = BundleTransferPolicy(max_inflight_put=8) + assert transfer._transport._resolve_buffer_group_put_mode(groups, policy) == "parallel" + + def test_put_object_roundtrips_numpy_and_torch_tensor_fields() -> None: torch = pytest.importorskip("torch") store, transfer = make_transfer() @@ -1603,9 +1622,9 @@ def test_dataproto_field_schema_encodes_typed_ragged_non_tensor_field() -> None: ) result = transfer.get_dataproto(ref)["non_tensor_batch"]["tokens"] - assert result[0] == [1, 2] + assert result[0].tolist() == [1, 2] assert result[1] is None - assert result[2] == [3] + assert result[2].tolist() == [3] bad_text = np.asarray([object()], dtype=object) with pytest.raises(AttributeError, match="failed to encode.*'text'.*utf8_ragged"): @@ -1814,6 +1833,59 @@ def test_unified_put_get_roundtrips_flat_dict() -> None: assert result["step"] == 7 +def test_unified_dict_get_returns_flat_lists_without_changing_dataproto_shape() -> None: + _store, transfer = make_transfer() + data = { + "tokens": [ + np.asarray([1, 2], dtype=np.int32), + np.asarray([3], dtype=np.int32), + np.asarray([4, 5, 6], dtype=np.int32), + ], + "json": [{"rank": 0}, {"rank": 1}, {"rank": 2}], + } + schemas = { + "tokens": FieldSchema( + codec="typed_ragged", + metadata={"section": "non_tensor_batch", "dtype": "int32"}, + ), + "json": FieldSchema( + codec="msgpack_ragged", + metadata={"section": "non_tensor_batch"}, + ), + } + + ref = transfer.put(data, type="dict", field_schemas=schemas) + + flat = transfer.get(ref, type="dict") + assert isinstance(flat["tokens"], list) + assert isinstance(flat["json"], list) + assert [row.tolist() for row in flat["tokens"]] == [[1, 2], [3], [4, 5, 6]] + assert flat["json"] == [{"rank": 0}, {"rank": 1}, {"rank": 2}] + + sliced = transfer.get(ref, type="dict", rows=slice(1, 3)) + assert isinstance(sliced["tokens"], list) + assert isinstance(sliced["json"], list) + assert [row.tolist() for row in sliced["tokens"]] == [[3], [4, 5, 6]] + assert sliced["json"] == [{"rank": 1}, {"rank": 2}] + + envelope = transfer.get_dataproto(ref, data_cls=dict) + non_tensor_batch = envelope["non_tensor_batch"] + assert isinstance(non_tensor_batch["tokens"], np.ndarray) + assert non_tensor_batch["tokens"].dtype == object + assert isinstance(non_tensor_batch["json"], np.ndarray) + assert non_tensor_batch["json"].dtype == object + assert [row.tolist() for row in non_tensor_batch["tokens"]] == [ + [1, 2], + [3], + [4, 5, 6], + ] + assert non_tensor_batch["json"].tolist() == [ + {"rank": 0}, + {"rank": 1}, + {"rank": 2}, + ] + + def test_unified_put_rejects_unknown_type() -> None: _store, transfer = make_transfer() @@ -1879,7 +1951,8 @@ def schema(codec: str, section: str, dtype: str | None = None) -> FieldSchema: result = transfer.get(ref, type="dict") assert np.array_equal(result["input_ids"], np.arange(2)) - assert result["tokens"] == [[1, 2], None] + assert result["tokens"][0].tolist() == [1, 2] + assert result["tokens"][1] is None assert result["partition"] == [0, 1] assert result["response_lengths"] == [2, 0] assert result["global_batch_sizes"] == [2] @@ -2899,7 +2972,11 @@ def test_dataproto_helper_typed_ragged_uses_multi_buffer_put() -> None: assert ref.encoded_non_tensor["tokens"]["codec"] == "typed_ragged" tokens = result["non_tensor_batch"]["tokens"].tolist() - assert tokens == [[1, 2], None, [3, 4, 5]] + assert [None if row is None else row.tolist() for row in tokens] == [ + [1, 2], + None, + [3, 4, 5], + ] assert rows[0].nbytes + rows[2].nbytes in pool.acquire_sizes @@ -2933,7 +3010,11 @@ def test_dataproto_helper_typed_ragged_fast_copy_put() -> None: result = transfer.get_dataproto(ref) tokens = result["non_tensor_batch"]["tokens"].tolist() - assert tokens == [[1, 2], None, [3, 4, 5]] + assert [None if row is None else row.tolist() for row in tokens] == [ + [1, 2], + None, + [3, 4, 5], + ] assert store.batch_put_from_calls > 0 assert rows[0].nbytes + rows[2].nbytes in pool.acquire_sizes @@ -2992,6 +3073,135 @@ def test_dataproto_helper_typed_ragged_zero_copy_rejects_source_buffers() -> Non ) +def test_typed_ragged_packs_ndarray_rows_contiguously() -> None: + _store, transfer = make_transfer() + rows = [ + np.asarray(7, dtype=np.int32), + np.arange(3, dtype=np.int32), + None, + np.arange(4, dtype=np.int32).reshape(2, 2), + np.arange(6, dtype=np.int32).reshape(2, 3), + ] + + ref = transfer.put( + {"values": rows}, + field_schemas={ + "values": FieldSchema( + codec="typed_ragged", + metadata={"section": "non_tensor_batch"}, + ) + }, + type="dict", + ) + result = transfer.get(ref, type="dict") + view = transfer.dataproto_manifest_view(ref) + + encoded = ref.encoded_non_tensor["values"] + assert encoded["codec"] == "typed_ragged" + assert encoded["metadata"]["physical_layout"] == "contiguous_flat" + assert view["non_tensor_fields"]["values"]["spec"]["payload_specs"]["data"][ + "shape" + ] == [14] + assert [ + None if row is None else row.tolist() for row in result["values"] + ] == [ + None if row is None else row.tolist() for row in rows + ] + assert all( + row is None or isinstance(row, np.ndarray) for row in result["values"] + ) + assert result["values"][0].shape == () + assert [row.dtype for row in result["values"] if row is not None] == [ + np.dtype(np.int32) + ] * 4 + gathered = transfer.get_dataproto( + ref, non_tensor_fields=["values"], rows=[4, 2, 0] + ) + assert [ + None if row is None else row.tolist() + for row in gathered["non_tensor_batch"]["values"] + ] == [ + rows[4].tolist(), + None, + rows[0].tolist(), + ] + assert gathered["non_tensor_batch"]["values"][2].shape == () + + +def _typed_ragged_payload_array( + payload: dict[str, object], dtype: np.dtype +) -> np.ndarray: + data = payload["data"] + arrays = getattr(data, "arrays", None) + if arrays is not None: + flat_arrays = [ + np.asarray(array, dtype=dtype).reshape(-1) for array in arrays + ] + if not flat_arrays: + return np.asarray([], dtype=dtype) + return np.concatenate(flat_arrays) + return np.concatenate([np.frombuffer(buf, dtype=dtype) for buf in data.buffers]) + +def test_typed_ragged_fast_decodes_equal_shape_ndarray_views() -> None: + rows = [np.arange(6, dtype=np.int32).reshape(2, 3) + i * 10 for i in range(3)] + payload, metadata = sos._encode_typed_ragged_values(rows, np.dtype(np.int32)) + data = _typed_ragged_payload_array(payload, np.dtype(np.int32)) + + decoded = sos._decode_typed_ragged_values({**payload, "data": data}, 3, metadata) + + assert metadata["physical_layout"] == "contiguous_flat" + assert all(isinstance(row, np.ndarray) for row in decoded) + assert all(np.shares_memory(row, data) for row in decoded) + assert [row.tolist() for row in decoded] == [row.tolist() for row in rows] + + +def test_typed_ragged_single_ndarray_row_uses_general_layout() -> None: + rows = [np.arange(6, dtype=np.int32).reshape(2, 3)] + + assert sos._encode_typed_ragged_regular_ndarray_rows( + rows, np.dtype(np.int32) + ) is None + payload, metadata = sos._encode_typed_ragged_values(rows, np.dtype(np.int32)) + data = _typed_ragged_payload_array(payload, np.dtype(np.int32)) + + decoded = sos._decode_typed_ragged_values({**payload, "data": data}, 1, metadata) + + assert metadata["physical_layout"] == "contiguous_flat" + assert isinstance(decoded[0], np.ndarray) + assert np.shares_memory(decoded[0], data) + assert decoded[0].tolist() == rows[0].tolist() + + +def test_typed_ragged_equal_shape_rows_with_nulls_use_general_layout() -> None: + rows = [ + np.arange(6, dtype=np.int32).reshape(2, 3), + None, + np.arange(6, 12, dtype=np.int32).reshape(2, 3), + ] + + assert sos._encode_typed_ragged_regular_ndarray_rows( + rows, np.dtype(np.int32) + ) is None + payload, metadata = sos._encode_typed_ragged_values(rows, np.dtype(np.int32)) + data = _typed_ragged_payload_array(payload, np.dtype(np.int32)) + + decoded = sos._decode_typed_ragged_values({**payload, "data": data}, 3, metadata) + + assert metadata["physical_layout"] == "contiguous_flat" + assert decoded[1] is None + assert all( + row is None or isinstance(row, np.ndarray) for row in decoded + ) + assert all( + row is None or np.shares_memory(row, data) for row in decoded + ) + assert [None if row is None else row.tolist() for row in decoded] == [ + rows[0].tolist(), + None, + rows[2].tolist(), + ] + + def test_dataproto_helper_rejects_unsupported_object_non_tensor() -> None: _store, transfer = make_transfer() data = SimpleDataProto( @@ -3026,6 +3236,7 @@ def test_dataproto_helper_ragged_tensor_non_tensor_roundtrip() -> None: result = transfer.get_dataproto(ref) assert ref.encoded_non_tensor["ragged"]["codec"] == "ragged_tensor" + assert ref.encoded_non_tensor["ragged"]["metadata"]["dtype"] == "torch.float32" actual = result["non_tensor_batch"]["ragged"] assert torch.equal(actual[0], ragged[0]) assert actual[1] is None @@ -3040,6 +3251,11 @@ def test_dataproto_helper_ragged_tensor_non_tensor_roundtrip() -> None: with pytest.raises(ValueError, match="mixed tensor dtype"): transfer.put_dataproto(SimpleDataProto(non_tensor_batch={"ragged": mixed})) + bfloat = np.empty(1, dtype=object) + bfloat[0] = torch.zeros(1, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="bfloat16"): + transfer.put_dataproto(SimpleDataProto(non_tensor_batch={"ragged": bfloat})) + def test_dataproto_helper_jagged_nested_batch_tensor_roundtrip() -> None: torch = pytest.importorskip("torch") @@ -3219,11 +3435,13 @@ def test_dataproto_helper_dict_of_native_object_leaves_uses_recursive_codec() -> assert ref.encoded_non_tensor["samples"]["codec"] == "structured_recursive" assert actual[0]["media"] == [b"a", b"bc"] assert actual[0]["blob"] == b"payload-0" - assert actual[0]["scores"] == [1, 2] + assert isinstance(actual[0]["scores"], np.ndarray) + assert actual[0]["scores"].tolist() == [1, 2] assert actual[0]["label"] == "x" assert actual[1]["media"] == [] assert actual[1]["blob"] == b"payload-1" - assert actual[1]["scores"] == [3] + assert isinstance(actual[1]["scores"], np.ndarray) + assert actual[1]["scores"].tolist() == [3] assert actual[1]["label"] == "y" assert actual[2] == {"label": "missing-native"} assert actual[3] is None From 9fbd622256e7a26aeea9f038e316eeb58203c102 Mon Sep 17 00:00:00 2001 From: enzodechine Date: Sat, 8 Aug 2026 22:55:12 +0800 Subject: [PATCH 006/483] [Bugfix] clean up partially initialized TENT RDMA contexts (#3325) * [Fix] clean up partially initialized TENT RDMA contexts * Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> * [Fix] prevent CQ leak on construction failure * [Fix] simplify cleanup of partially initialized TENT RDMA contexts * [Fix] hold mr_set_mutex_ and correct cleanup comment in RdmaContext --------- Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- .../include/tent/transport/rdma/context.h | 6 +++ .../tent/src/transport/rdma/context.cpp | 44 +++++++++++-------- .../src/transport/rdma/rdma_transport.cpp | 4 +- 3 files changed, 35 insertions(+), 19 deletions(-) diff --git a/mooncake-transfer-engine/tent/include/tent/transport/rdma/context.h b/mooncake-transfer-engine/tent/include/tent/transport/rdma/context.h index 2c010872e6..8156ba808a 100644 --- a/mooncake-transfer-engine/tent/include/tent/transport/rdma/context.h +++ b/mooncake-transfer-engine/tent/include/tent/transport/rdma/context.h @@ -24,6 +24,7 @@ #include #include #include +#include #include #include #include @@ -123,6 +124,11 @@ class RdmaContext { private: int openDevice(const std::string &device_name, uint8_t port); + // Release every resource currently owned by this context. This is + // intentionally state-independent so it can clean up a partially completed + // enable() and can safely be retried. + void cleanupResources(); + private: // initialized during ctor, will never be changed during the context's // lifecycle diff --git a/mooncake-transfer-engine/tent/src/transport/rdma/context.cpp b/mooncake-transfer-engine/tent/src/transport/rdma/context.cpp index d96081dc6f..ca4629eea8 100644 --- a/mooncake-transfer-engine/tent/src/transport/rdma/context.cpp +++ b/mooncake-transfer-engine/tent/src/transport/rdma/context.cpp @@ -301,9 +301,7 @@ RdmaContext::RdmaContext(RdmaTransport& transport) std::call_once(g_once_flag, fork_init); } -RdmaContext::~RdmaContext() { - if (status_ != DEVICE_UNINIT) disable(); -} +RdmaContext::~RdmaContext() { disable(); } int RdmaContext::construct(const std::string& device_name, std::shared_ptr params) { @@ -320,10 +318,15 @@ int RdmaContext::construct(const std::string& device_name, } int RdmaContext::enable() { - if (status_ != DEVICE_DISABLED) { + if (status_ == DEVICE_ENABLED || status_ == DEVICE_PAUSED) { LOG(WARNING) << "RDMA context " << name() << " has been enabled"; return 0; } + if (status_ != DEVICE_DISABLED) { + LOG(ERROR) << "RDMA context " << name() << " cannot be enabled from " + << statusToString(status_); + return -1; + } if (openDevice(device_name_, params_->device.port)) { LOG(ERROR) << "Failed to open device [" << device_name_ << "] on port [" << params_->device.port << "] with GID index [" @@ -369,13 +372,13 @@ int RdmaContext::enable() { } for (int i = 0; i < params_->device.num_cq_list; ++i) { - auto cq = new RdmaCQ(); + auto cq = std::make_unique(); int ret = cq->construct(this, params_->device.max_cqe, i); if (ret) { disable(); return ret; } - cq_list_.push_back(cq); + cq_list_.push_back(cq.release()); } // Create dedicated notification CQ @@ -419,9 +422,7 @@ int RdmaContext::enable() { if (ret) { PLOG(ERROR) << "Failed to query port " << params_->device.port << " on " << device_name_; - if (verbs_.ibv_close_device(native_context_)) { - PLOG(ERROR) << "ibv_close_device"; - } + disable(); return -1; } @@ -439,21 +440,31 @@ int RdmaContext::enable() { } int RdmaContext::disable() { - if (status_ == DEVICE_UNINIT || status_ == DEVICE_DISABLED) { + if (!native_context_ && !native_pd_ && event_fd_ < 0 && + comp_channel_.empty() && cq_list_.empty() && !notify_cq_) { LOG(WARNING) << "RDMA context " << name() << " has been deconstructed"; return 0; } - if (endpoint_store_->clear()) { + if (endpoint_store_ && endpoint_store_->clear()) { LOG(ERROR) << "Failed to destroy all endpoints for context " << name() << "; preserving CQ, PD and device resources for retry"; return -1; } - for (auto& entry : mr_set_) { - int ret = verbs_.ibv_dereg_mr(entry); - if (ret) PLOG(ERROR) << "ibv_dereg_mr"; + cleanupResources(); + status_ = DEVICE_DISABLED; + return 0; +} + +void RdmaContext::cleanupResources() { + { + std::lock_guard lock(mr_set_mutex_); + for (auto& entry : mr_set_) { + int ret = verbs_.ibv_dereg_mr(entry); + if (ret) PLOG(ERROR) << "ibv_dereg_mr"; + } + mr_set_.clear(); } - mr_set_.clear(); for (auto& entry : cq_list_) { delete entry; } @@ -486,9 +497,6 @@ int RdmaContext::disable() { PLOG(ERROR) << "ibv_close_device"; native_context_ = nullptr; } - - status_ = DEVICE_DISABLED; - return 0; } int RdmaContext::pause() { diff --git a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp index 6a82670d45..b5563f7866 100644 --- a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp @@ -654,7 +654,9 @@ int RdmaTransport::onSetupRdmaConnections(const BootstrapDesc& peer_desc, } auto index = context_name_lookup_[local_nic_name]; auto context = context_set_[index]; - if (context->status() == RdmaContext::DEVICE_DISABLED) { + auto ctx_status = context->status(); + if (ctx_status != RdmaContext::DEVICE_ENABLED && + ctx_status != RdmaContext::DEVICE_PAUSED) { std::stringstream ss; ss << "Device is down: " << peer_desc.local_nic_path; LOG(ERROR) << ss.str(); From 7147f1e4f463cd0928b0923279ca9fd4948c364d Mon Sep 17 00:00:00 2001 From: SongOf <46475785+SongOf@users.noreply.github.com> Date: Sun, 9 Aug 2026 14:39:01 +0800 Subject: [PATCH 007/483] [TENT] Keep a never-constructed context in the NicID slot (#3339) Co-authored-by: maxlisongsong --- .../src/transport/rdma/rdma_transport.cpp | 22 +++++++++++-------- .../tent/tests/rdma_transport_test.cpp | 15 ++++++++----- 2 files changed, 22 insertions(+), 15 deletions(-) diff --git a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp index b5563f7866..f4d2e6177b 100644 --- a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp @@ -278,18 +278,22 @@ size_t RdmaTransport::initializeContexts() { size_t context_count = 0; for (size_t i = 0; i < local_topology_->getNicCount(); ++i) { auto entry = local_topology_->getNicEntry(i); - auto context = std::make_shared(*this); - context_set_.push_back(context); - if (entry->type != Topology::NIC_RDMA) continue; - int ret = context->construct(entry->name, params_); - if (ret) { + if (entry->type == Topology::NIC_RDMA) { + auto context = std::make_shared(*this); + if (context->construct(entry->name, params_) == 0) { + context_name_lookup_[entry->name] = i; + ++context_count; + local_buffer_manager_.addDevice(context.get()); + context_set_.push_back(std::move(context)); + continue; + } LOG(WARNING) << "Disable RDMA device " << entry->name << " because " << "of initialization failure"; - continue; } - context_name_lookup_[entry->name] = i; - ++context_count; - local_buffer_manager_.addDevice(context.get()); + // A never-constructed context, not the one whose construct() failed: + // the slot only has to stand in for the NicID, so it should not carry + // a device name or an endpoint store it will never use. + context_set_.push_back(std::make_shared(*this)); } return context_count; } diff --git a/mooncake-transfer-engine/tent/tests/rdma_transport_test.cpp b/mooncake-transfer-engine/tent/tests/rdma_transport_test.cpp index 6e918a21d2..583e8b0b6d 100644 --- a/mooncake-transfer-engine/tent/tests/rdma_transport_test.cpp +++ b/mooncake-transfer-engine/tent/tests/rdma_transport_test.cpp @@ -31,7 +31,6 @@ #include "tent/transfer_engine.h" #include "tent/runtime/topology.h" #include "tent/transport/rdma/context.h" -#include "tent/transport/rdma/ibv_loader.h" #include "tent/transport/rdma/params.h" #include "tent/transport/rdma/rdma_transport.h" #include "tent/transport/rdma/workers.h" @@ -165,6 +164,13 @@ void expectInertContextPerNic(const RdmaContextSet& contexts, size_t expected) { // Inert contexts must stay safe for the whole-list consumers. EXPECT_EQ(contexts[i]->cq(0), nullptr); EXPECT_EQ(contexts[i]->notifyCq(), nullptr); + EXPECT_LT(contexts[i]->eventFd(), 0); + EXPECT_EQ(contexts[i]->nativeContext(), nullptr); + // construct() sets device_name_ before it can fail, so an empty name + // is what tells a fresh placeholder from a kept half-built context. + EXPECT_TRUE(contexts[i]->name().empty()) + << "slot " << i << " kept a half-built context for " + << contexts[i]->name(); } } @@ -213,12 +219,9 @@ TEST(RdmaNicIndexAlignmentTest, ReclaimTickSkipsInertContexts) { RdmaTransportTestPeer::reclaimEndpoints(transport); } -// The construct()-failure branch. IbvLoader dlcloses libibverbs when no device -// is present, so construct() cannot be driven safely in that state. +// The construct()-failure branch. Runs on any host: this binary links +// libibverbs directly, so IbvLoader's dlclose does not unmap it. TEST(RdmaNicIndexAlignmentTest, ContextSetKeepsOneSlotWhenConstructFails) { - if (!IbvLoader::Instance().available()) - GTEST_SKIP() << "no usable libibverbs; construct() cannot be driven"; - // Device names that resolve to no real RNIC, so construct() fails on any // host, with a non-RDMA entry in the middle to offset the indexes. auto topology = std::make_shared(); From 4f1fedfb58c32d16ce0ad73089d4f07200ba1be3 Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Sun, 9 Aug 2026 15:27:24 +0800 Subject: [PATCH 008/483] [Store] Support pinned SSD-to-GPU restore (#3319) --- docs/source/deployment/ssd/ssd-offload.md | 3 + docs/source/design/ssd-offload.md | 12 +- mooncake-store/include/client_service.h | 19 ++- mooncake-store/include/file_storage.h | 28 +++- mooncake-store/include/pinned_buffer_pool.h | 20 ++- mooncake-store/include/real_client.h | 13 +- mooncake-store/include/storage_backend.h | 2 + mooncake-store/include/transfer_task.h | 18 ++- mooncake-store/src/client_service.cpp | 13 +- mooncake-store/src/file_storage.cpp | 89 +++++++++-- mooncake-store/src/real_client.cpp | 135 +++++++++++++---- mooncake-store/src/transfer_task.cpp | 158 +++++++++----------- mooncake-store/tests/CMakeLists.txt | 20 ++- mooncake-store/tests/file_storage_test.cpp | 59 ++++++-- mooncake-store/tests/pybind_client_test.cpp | 136 +++++++++++++++++ mooncake-store/tests/transfer_task_test.cpp | 96 +++++++++++- 16 files changed, 637 insertions(+), 184 deletions(-) diff --git a/docs/source/deployment/ssd/ssd-offload.md b/docs/source/deployment/ssd/ssd-offload.md index 34d2dd28a9..e4b59b8f6c 100644 --- a/docs/source/deployment/ssd/ssd-offload.md +++ b/docs/source/deployment/ssd/ssd-offload.md @@ -133,6 +133,7 @@ Start with `--enable_offload=true` for eager SSD persistence. Add `--offload_on_ | `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` | `/data/file_storage` | Absolute path to the SSD storage directory | | `MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR` | `bucket_storage_backend` | Storage backend type (see below) | | `MOONCAKE_OFFLOAD_LOCAL_BUFFER_SIZE_BYTES` | `1342177280` (1.25 GB) | Client-side staging buffer size | +| `MC_STORE_PINNED_RESTORE_ARENA_SIZE_BYTES` | `0` | Size of the additional preallocated pinned-host arena for same-process SSD-to-GPU restores. See the constraints below | | `MOONCAKE_OFFLOAD_SCANMETA_ITERATOR_KEYS_LIMIT` | `20000` | Max keys processed per iteration when scanning existing SSD metadata on startup | | `MOONCAKE_OFFLOAD_TOTAL_SIZE_LIMIT_BYTES` | `2199023255552` (2 TB) | Maximum disk usage | | `MOONCAKE_OFFLOAD_TOTAL_KEYS_LIMIT` | `10000000` | Maximum number of objects on disk | @@ -146,6 +147,8 @@ Start with `--enable_offload=true` for eager SSD persistence. Add `--offload_on_ The `MOONCAKE_OFFLOAD_*` watermark names are preferred. Short aliases `MOONCAKE_DISK_EVICTION_HIGH_WATERMARK_RATIO` and `MOONCAKE_DISK_EVICTION_LOW_WATERMARK_RATIO` are also accepted. The high watermark must be greater than the low watermark. +The pinned restore quota is separate from `MOONCAKE_OFFLOAD_LOCAL_BUFFER_SIZE_BYTES`; it does not convert the normal FileStorage arena to pinned memory. It is allocated only when `MC_STORE_MEMCPY=1`, and selected only when the current process owns the SSD replica and the restore destination is GPU memory. Tensor payload ranges are copied from their source offset without an additional FileStorage staging allocation. Remote, CPU-destination, and io_uring reads continue to use the existing path. If the pinned quota is exhausted, the request uses the normal restore arena; if the quota cannot be pinned at startup, the optimization remains disabled. The file-per-key backend may still use its own temporary pageable buffer internally; the bucket and offset-allocator backends can read into the supplied restore buffer directly. + ### Bucket backend settings Applies when `MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR=bucket_storage_backend`. diff --git a/docs/source/design/ssd-offload.md b/docs/source/design/ssd-offload.md index 1ca2607499..71cd369ded 100644 --- a/docs/source/design/ssd-offload.md +++ b/docs/source/design/ssd-offload.md @@ -93,7 +93,7 @@ Step by step: ### Load (SSD → memory) -The load path involves three parties: the **requesting client**, the **target client** that holds the SSD data, and the **Transfer Engine** for zero-copy data movement. +The default load path involves three parties: the **requesting client**, the **target client** that holds the SSD data, and the **Transfer Engine** for data movement. ``` Requesting Client Target Client Master @@ -119,13 +119,15 @@ Requesting Client Target Client Master │ │ (free ClientBuffer slot)│ ``` +For a same-process GPU destination, FileStorage reads into a quota-bounded pinned restore arena and submits H2D copies to the caller's GPU slices. Tensor ranges use a source offset, avoiding another CPU staging allocation. This path requires local memcpy and a pinned quota; CPU and remote restores are unchanged. + Step by step: 1. **Query master**: The requesting client calls `client_->BatchGet(keys, ...)` to query the master for replica locations. If the object has been offloaded, the master returns a `LOCAL_DISK` replica descriptor containing the target client's RPC address (`transport_endpoint`). -2. **RPC to target client**: The requesting client calls `batch_get_offload_object(keys, sizes)` on the target client identified by `transport_endpoint`. The target client calls `FileStorage::BatchGet`, which allocates slots in `ClientBuffer` and reads the requested objects from SSD via `StorageBackend::BatchLoad`. -3. **Response with buffer pointers**: The target client returns a `BatchGetOffloadObjectResponse` containing `batch_id`, a list of buffer `pointers` (addresses within `ClientBuffer`), the Transfer Engine address, and `gc_ttl_ms` (the buffer lease TTL). -4. **Zero-copy transfer**: The requesting client invokes `client_->BatchGetOffloadObject(transfer_engine_addr, keys, pointers, slices)`, which uses the Transfer Engine (RDMA or TCP) to pull the data directly from the target client's `ClientBuffer` into the application's target memory (DRAM or VRAM). No intermediate copy is made on the requesting client side. -5. **Release buffer**: After the transfer completes, the requesting client immediately calls `release_offload_buffer(batch_id)` on the target client to free the `ClientBuffer` slots. If the transfer takes longer than `gc_ttl_ms`, the buffer GC thread reclaims the slot automatically as a fallback. +2. **Select the owner path**: For a different process, the requesting client calls `batch_get_offload_object(keys, sizes)` on the target client identified by `transport_endpoint`. For a same-process GPU restore with pinned quota and local memcpy enabled, it calls the local `FileStorage::BatchGetLocal` path. +3. **Response with buffer pointers**: Remote requests receive pointers within `ClientBuffer`, plus a `batch_id`, Transfer Engine address, and buffer lease TTL. The local path receives pointers within the pinned restore arena and keeps their allocation owner in the requesting process. +4. **Transfer data**: Remote restores use the Transfer Engine (RDMA or TCP). The same-process pinned branch submits H2D copies from the local restore arena and supports per-object source offsets for tensor payloads. +5. **Release buffer**: Remote restores call `release_offload_buffer(batch_id)` and retain TTL GC as a failure fallback. A same-process allocation is held by an RAII owner through the synchronous H2D operation and released automatically on success or failure; it is never published in the remote batch map. --- diff --git a/mooncake-store/include/client_service.h b/mooncake-store/include/client_service.h index e89f911e6f..7bd2e481de 100644 --- a/mooncake-store/include/client_service.h +++ b/mooncake-store/include/client_service.h @@ -495,13 +495,6 @@ class Client { /** * @brief Performs a batched read of multiple objects using a * high-throughput Transfer Engine. - * @param transfer_engine_addr Address of the Transfer Engine service (e.g., - * "ip:port"). - * @param keys List of keys identifying the data objects to be transferred - * @param pointers Array of destination memory addresses on the remote node - * where data will be written (one per key) - * @param batch_slices Map from object key to its data slice - * (`mooncake::Slice`), containing raw bytes to be written. */ tl::expected BatchGetOffloadObject( const std::string& transfer_engine_addr, @@ -510,6 +503,13 @@ class Client { const std::unordered_map>& batch_slices); + tl::expected BatchGetOffloadObject( + const std::string& transfer_engine_addr, + const std::vector& keys, + const std::vector& pointers, + const std::unordered_map>& batch_slices, + OffloadBufferAccess buffer_access); + /** * @brief Notifies the master that offloading of specified objects has * succeeded. @@ -584,6 +584,11 @@ class Client { [[nodiscard]] const std::string& GetProtocol() const { return protocol_; } + [[nodiscard]] bool CanUseLocalMemcpy(const std::string& endpoint) const { + return transfer_submitter_ != nullptr && + transfer_submitter_->canUseLocalMemcpy(endpoint); + } + /** * @brief Get the endpoint address for segment operations. * @return For P2PHANDSHAKE mode, returns the actual RPC endpoint (IP:Port). diff --git a/mooncake-store/include/file_storage.h b/mooncake-store/include/file_storage.h index 501a80c66f..174fb5fc1f 100644 --- a/mooncake-store/include/file_storage.h +++ b/mooncake-store/include/file_storage.h @@ -29,6 +29,14 @@ class FileStorage { std::vector pointers; }; + struct LocalBatchResult { + std::vector pointers; + + private: + friend class FileStorage; + std::shared_ptr owner; + }; + /** * @brief Reads multiple key-value (KV) entries from local storage and * forwards them to a remote node. @@ -41,6 +49,14 @@ class FileStorage { const std::vector& keys, const std::vector& sizes); + tl::expected BatchGetLocal( + const std::vector& keys, + const std::vector& sizes); + + [[nodiscard]] bool HasPinnedRestoreArena() const { + return pinned_restore_arena_allocator_ != nullptr; + } + FileStorageConfig config_; /** @@ -142,8 +158,12 @@ class FileStorage { tl::expected RegisterLocalMemory(); tl::expected, ErrorCode> AllocateBatch( - const std::vector& keys, - const std::vector& sizes); + const std::vector& keys, const std::vector& sizes, + ClientBufferAllocator& allocator); + + tl::expected, ErrorCode> LoadBatch( + const std::vector& keys, const std::vector& sizes, + bool prefer_pinned); void ClientBufferGCThreadFunc(); @@ -157,8 +177,10 @@ class FileStorage { std::shared_ptr client_; SsdMetric* ssd_metric_{nullptr}; std::string local_rpc_addr_; - // Pinned host memory pool for GPU D2H staging in OffloadObjects + // Pinned memory for GPU staging and SSD-to-GPU restores. std::unique_ptr pinned_buffer_pool_; + PinnedBufferPool::Buffer pinned_restore_arena_; + std::shared_ptr pinned_restore_arena_allocator_; std::shared_ptr storage_backend_; std::shared_ptr client_buffer_allocator_; mutable Mutex client_buffer_mutex_; diff --git a/mooncake-store/include/pinned_buffer_pool.h b/mooncake-store/include/pinned_buffer_pool.h index 25461f7d1f..5f50dde1d5 100644 --- a/mooncake-store/include/pinned_buffer_pool.h +++ b/mooncake-store/include/pinned_buffer_pool.h @@ -80,7 +80,6 @@ class PinnedBufferPool { for (size_t i = 0; i < pool_.size(); ++i) { if (pool_[i].capacity >= size) { Buffer buf = std::move(pool_[i]); - // O(1) erase: swap with back then pop if (i != pool_.size() - 1) { pool_[i] = std::move(pool_.back()); } @@ -92,6 +91,17 @@ class PinnedBufferPool { return AllocNew(size); } + // Never falls back to pageable memory. + static Buffer AllocatePinned(size_t size) { + auto accelerators = + device::GetAcceleratorRegistry().RuntimeAccelerators(); + for (auto* accelerator : accelerators.Devices()) { + auto host = accelerator->AllocatePinnedHost(size); + if (host.addr) return Buffer(std::move(host)); + } + return {}; + } + void Release(Buffer buf) { std::lock_guard lk(mutex_); if (pool_.size() < max_pool_size_) { @@ -112,12 +122,8 @@ class PinnedBufferPool { private: static Buffer AllocNew(size_t size) { - const auto& registry = device::GetAcceleratorRegistry(); - auto runtime_accelerator = registry.RuntimeAccelerators(); - for (auto* accelerator : runtime_accelerator.Devices()) { - auto host = accelerator->AllocatePinnedHost(size); - if (host.addr) return Buffer(std::move(host)); - } + auto pinned = AllocatePinned(size); + if (pinned.data) return pinned; return Buffer::Pageable(size); } diff --git a/mooncake-store/include/real_client.h b/mooncake-store/include/real_client.h index 89bf4a4c4e..b0f9f20006 100644 --- a/mooncake-store/include/real_client.h +++ b/mooncake-store/include/real_client.h @@ -715,9 +715,20 @@ class RealClient : public PyClient { "ip:port"). */ + struct OffloadReadRange { + uint64_t source_offset; + int64_t restore_size; + }; + tl::expected batch_get_into_offload_object_internal( const std::string &target_rpc_service_addr, - std::unordered_map> &objects); + std::unordered_map> &objects, + const OffloadReadRange *read_range = nullptr); + + bool can_use_pinned_restore_arena( + const std::string &target_rpc_service_addr, + const std::unordered_map> &objects) + const; int64_t get_offload_rpc_read_count() const { return offload_rpc_read_count_.load(std::memory_order_relaxed); diff --git a/mooncake-store/include/storage_backend.h b/mooncake-store/include/storage_backend.h index 83afc3c700..14e8a822af 100644 --- a/mooncake-store/include/storage_backend.h +++ b/mooncake-store/include/storage_backend.h @@ -300,6 +300,8 @@ struct FileStorageConfig { // Size of the local client-side buffer (used for caching or batching) int64_t local_buffer_size = 1280 * kMB; // ~1.2 GB + int64_t pinned_restore_arena_size = 0; + // Limits for scanning and iteration operations int64_t scanmeta_iterator_keys_limit = 20000; // Max number of keys returned per Scan call, required by bucket diff --git a/mooncake-store/include/transfer_task.h b/mooncake-store/include/transfer_task.h index 5fe938e09e..8676cd64b8 100644 --- a/mooncake-store/include/transfer_task.h +++ b/mooncake-store/include/transfer_task.h @@ -37,6 +37,11 @@ enum class TransferStrategy { SPDK_NVMF = 4 // Spdk nvmf operation }; +enum class OffloadBufferAccess { + kTransferEngine, + kLocalAddress, +}; + /** * @brief Stream operator for TransferStrategy */ @@ -573,7 +578,10 @@ class TransferSubmitter { const std::vector& keys, const std::vector& pointers, const std::unordered_map>& - batched_slices); + batched_slices, + OffloadBufferAccess buffer_access); + + [[nodiscard]] bool canUseLocalMemcpy(const std::string& endpoint) const; /** * @brief Pure comparison helper: returns true iff both endpoints are @@ -608,11 +616,6 @@ class TransferSubmitter { TransferStrategy selectStrategy(const AllocatedBuffer::Descriptor& handle, const std::vector& slices) const; - /** - * @brief Check if all handles refer to local segments - */ - bool isLocalTransfer(const AllocatedBuffer::Descriptor& handle) const; - /** * @brief Validate transfer parameters */ @@ -627,6 +630,9 @@ class TransferSubmitter { const std::vector& slices, const TransferRequest::OpCode op_code, uint64_t src_offset = 0); + std::optional submitMemcpyOperations( + std::vector operations); + #ifdef USE_NOF /** * @brief Submit SPDK NVMe-oF operation asynchronously diff --git a/mooncake-store/src/client_service.cpp b/mooncake-store/src/client_service.cpp index 137d4b434c..2ec06bf52b 100644 --- a/mooncake-store/src/client_service.cpp +++ b/mooncake-store/src/client_service.cpp @@ -3314,8 +3314,19 @@ tl::expected Client::BatchGetOffloadObject( const std::vector& keys, const std::vector& pointers, const std::unordered_map>& batch_slices) { + return BatchGetOffloadObject(transfer_engine_addr, keys, pointers, + batch_slices, + OffloadBufferAccess::kTransferEngine); +} + +tl::expected Client::BatchGetOffloadObject( + const std::string& transfer_engine_addr, + const std::vector& keys, + const std::vector& pointers, + const std::unordered_map>& batch_slices, + OffloadBufferAccess buffer_access) { auto future = transfer_submitter_->submit_batch_get_offload_object( - transfer_engine_addr, keys, pointers, batch_slices); + transfer_engine_addr, keys, pointers, batch_slices, buffer_access); if (!future) { LOG(ERROR) << "Failed to submit transfer operation"; return tl::make_unexpected(ErrorCode::TRANSFER_FAIL); diff --git a/mooncake-store/src/file_storage.cpp b/mooncake-store/src/file_storage.cpp index f4ccfbc0a0..09ae610d5c 100644 --- a/mooncake-store/src/file_storage.cpp +++ b/mooncake-store/src/file_storage.cpp @@ -108,6 +108,10 @@ FileStorageConfig FileStorageConfig::FromEnvironment() { config.local_buffer_size = Environ::GetInt64( "MOONCAKE_OFFLOAD_LOCAL_BUFFER_SIZE_BYTES", config.local_buffer_size); + config.pinned_restore_arena_size = + Environ::GetInt64("MC_STORE_PINNED_RESTORE_ARENA_SIZE_BYTES", + config.pinned_restore_arena_size); + config.scanmeta_iterator_keys_limit = Environ::GetInt64( "MOONCAKE_OFFLOAD_SCANMETA_ITERATOR_KEYS_LIMIT", Environ::GetInt64("MOONCAKE_SCANMETA_ITERATOR_KEYS_LIMIT", @@ -226,6 +230,11 @@ bool FileStorageConfig::Validate() const { LOG(ERROR) << "FileStorageConfig: total_size_limit should not be zero"; return false; } + if (pinned_restore_arena_size < 0) { + LOG(ERROR) << "FileStorageConfig: pinned_restore_arena_size must be " + "non-negative"; + return false; + } if (heartbeat_interval_seconds <= 0) { LOG(ERROR) << "FileStorageConfig: heartbeat_interval_seconds must > 0"; return false; @@ -267,6 +276,29 @@ FileStorage::FileStorage(const FileStorageConfig& config, throw std::invalid_argument("Invalid FileStorage configuration"); } + if (config.pinned_restore_arena_size > 0) { + if (config.use_uring) { + LOG(WARNING) << "Pinned SSD restore is disabled with io_uring"; + } else if (!client || + !client->CanUseLocalMemcpy(client->GetSegmentEndpoint())) { + LOG(WARNING) + << "Pinned SSD restore is disabled: local memcpy unavailable"; + } else { + auto buffer = PinnedBufferPool::AllocatePinned( + static_cast(config.pinned_restore_arena_size)); + if (buffer.pinned_host.addr) { + pinned_restore_arena_ = std::move(buffer); + pinned_restore_arena_allocator_ = ClientBufferAllocator::create( + pinned_restore_arena_.data, pinned_restore_arena_.capacity, + client->GetProtocol()); + LOG(INFO) << "Initialized pinned SSD restore arena, size=" + << pinned_restore_arena_.capacity; + } else { + LOG(WARNING) << "Failed to allocate pinned SSD restore arena"; + } + } + } + auto create_storage_backend_result = CreateStorageBackend(config_); if (!create_storage_backend_result) { LOG(ERROR) << "Failed to create storage backend"; @@ -398,15 +430,23 @@ tl::expected FileStorage::Init() { return {}; } -tl::expected FileStorage::BatchGet( - const std::vector& keys, const std::vector& sizes) { - auto start_time = std::chrono::steady_clock::now(); - auto allocate_res = AllocateBatch(keys, sizes); +tl::expected, ErrorCode> +FileStorage::LoadBatch(const std::vector& keys, + const std::vector& sizes, bool prefer_pinned) { + const bool use_pinned = prefer_pinned && pinned_restore_arena_allocator_; + auto& allocator = use_pinned ? *pinned_restore_arena_allocator_ + : *client_buffer_allocator_; + auto allocate_res = AllocateBatch(keys, sizes, allocator); + if (!allocate_res && use_pinned && + allocate_res.error() == ErrorCode::BUFFER_OVERFLOW) { + VLOG(1) << "Pinned SSD restore arena exhausted; using default arena"; + allocate_res = AllocateBatch(keys, sizes, *client_buffer_allocator_); + } if (!allocate_res) { LOG(ERROR) << "Failed to allocate batch objects"; return tl::make_unexpected(allocate_res.error()); } - auto allocated_batch = allocate_res.value(); + auto allocated_batch = std::move(allocate_res.value()); auto result = BatchLoad(allocated_batch->slices); if (!result) { LOG(ERROR) << "Batch load object failed,err_code = " << result.error(); @@ -424,6 +464,16 @@ tl::expected FileStorage::BatchGet( } } + return allocated_batch; +} + +tl::expected FileStorage::BatchGet( + const std::vector& keys, const std::vector& sizes) { + auto start_time = std::chrono::steady_clock::now(); + auto load_result = LoadBatch(keys, sizes, false); + if (!load_result) return tl::make_unexpected(load_result.error()); + + auto allocated_batch = std::move(load_result.value()); uint64_t batch_id = allocated_batch->batch_id; BatchGetResult batch_result{batch_id, allocated_batch->pointers}; @@ -439,6 +489,19 @@ tl::expected FileStorage::BatchGet( return batch_result; } +tl::expected +FileStorage::BatchGetLocal(const std::vector& keys, + const std::vector& sizes) { + auto load_result = LoadBatch(keys, sizes, true); + if (!load_result) return tl::make_unexpected(load_result.error()); + + auto batch = std::move(load_result.value()); + LocalBatchResult result; + result.pointers = std::move(batch->pointers); + result.owner = std::move(batch); + return result; +} + bool FileStorage::IsPerBucketSoftOffloadError(ErrorCode error) { return error == ErrorCode::INVALID_READ || error == ErrorCode::OBJECT_ALREADY_EXISTS; @@ -1023,7 +1086,8 @@ tl::expected FileStorage::ProcessPromotionTasks() { // staging space when the local goes out of scope. std::vector single_key{storage_key}; std::vector single_size{size}; - auto allocate_res = AllocateBatch(single_key, single_size); + auto allocate_res = + AllocateBatch(single_key, single_size, *client_buffer_allocator_); if (!allocate_res) { LOG(WARNING) << "Promotion: AllocateBatch failed for key=" << key << ", error=" << allocate_res.error(); @@ -1156,7 +1220,8 @@ tl::expected FileStorage::RegisterLocalMemory() { tl::expected, ErrorCode> FileStorage::AllocateBatch(const std::vector& keys, - const std::vector& sizes) { + const std::vector& sizes, + ClientBufferAllocator& allocator) { if (keys.size() != sizes.size()) { LOG(ERROR) << "Mismatched keys and sizes count: keys=" << keys.size() << ", sizes=" << sizes.size(); @@ -1187,8 +1252,9 @@ FileStorage::AllocateBatch(const std::vector& keys, size_t alloc_size = align_up(data_size, kDirectIOAlignment) + 2 * kDirectIOAlignment; - auto alloc_result = client_buffer_allocator_->allocate(alloc_size); - if (!alloc_result && !gc_triggered) { + auto alloc_result = allocator.allocate(alloc_size); + if (!alloc_result && !gc_triggered && + &allocator == client_buffer_allocator_.get()) { gc_triggered = true; { MutexLocker locker(&client_buffer_mutex_); @@ -1202,12 +1268,9 @@ FileStorage::AllocateBatch(const std::vector& keys, } } } - alloc_result = client_buffer_allocator_->allocate(alloc_size); + alloc_result = allocator.allocate(alloc_size); } if (!alloc_result) { - LOG(ERROR) << "Failed to allocate slice buffer, size = " - << alloc_size << " (data_size=" << data_size - << "), key = " << keys[i]; return tl::make_unexpected(ErrorCode::BUFFER_OVERFLOW); } diff --git a/mooncake-store/src/real_client.cpp b/mooncake-store/src/real_client.cpp index 3f5eb35fbd..4789f789ca 100644 --- a/mooncake-store/src/real_client.cpp +++ b/mooncake-store/src/real_client.cpp @@ -3431,12 +3431,27 @@ tl::expected RealClient::execute_ranged_read( }; if (replica.is_local_disk_replica()) { + const auto &endpoint = + replica.get_local_disk_descriptor().transport_endpoint; + void *dst = static_cast(buffer) + dst_offset; + std::unordered_map> objects{ + {key, {{dst, size}}}}; + if (can_use_pinned_restore_arena(endpoint, objects)) { + if (total_size > uint64_t(std::numeric_limits::max())) { + return tl::unexpected(ErrorCode::INVALID_PARAMS); + } + const OffloadReadRange read_range{src_offset, + static_cast(total_size)}; + auto result = batch_get_into_offload_object_internal( + endpoint, objects, &read_range); + if (!result) return tl::unexpected(result.error()); + return static_cast(size); + } + // LOCAL_DISK: offload RPC transfers sequentially from remote offset // 0, so we only need src_offset + size bytes (not total_size). return partial_disk_read( [&](void *tmp_buf) -> tl::expected { - const auto &endpoint = - replica.get_local_disk_descriptor().transport_endpoint; std::unordered_map> objects; objects.emplace( key, std::vector{{static_cast(tmp_buf), @@ -5920,38 +5935,113 @@ bool RealClient::release_offload_buffer(uint64_t batch_id) { return file_storage_->ReleaseBuffer(batch_id); } +bool RealClient::can_use_pinned_restore_arena( + const std::string &target_rpc_service_addr, + const std::unordered_map> &objects) const { + if (!file_storage_ || target_rpc_service_addr != local_rpc_addr || + !file_storage_->HasPinnedRestoreArena()) { + return false; + } + auto accelerators = device::GetAcceleratorRegistry().RuntimeAccelerators(); + bool has_data = false; + for (const auto &object : objects) { + for (const auto &slice : object.second) { + if (slice.size && !accelerators.FindDeviceForPointer(slice.ptr)) { + return false; + } + has_data |= slice.size != 0; + } + } + return has_data; +} + tl::expected RealClient::batch_get_into_offload_object_internal( const std::string &target_rpc_service_addr, - std::unordered_map> &objects) { + std::unordered_map> &objects, + const OffloadReadRange *read_range) { offload_rpc_read_count_.fetch_add(1, std::memory_order_relaxed); auto start_time = std::chrono::steady_clock::now(); std::vector keys; std::vector storage_keys; std::vector sizes; + if (read_range && objects.size() != 1) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } const TenantId tenant_id(client_->tenant_id()); for (const auto &object_it : objects) { keys.emplace_back(object_it.first); storage_keys.emplace_back(tenant_id.MakeScopedKey(object_it.first)); - int64_t total = 0; - for (const auto &s : object_it.second) total += s.size; - sizes.emplace_back(total); - } - auto batchGetResp = client_requester_->batch_get_offload_object( - target_rpc_service_addr, storage_keys, sizes); - if (!batchGetResp) { + uint64_t total = 0; + for (const auto &slice : object_it.second) { + if (slice.size > std::numeric_limits::max() - total) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + total += slice.size; + } + if (total > uint64_t(std::numeric_limits::max())) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + int64_t storage_size = static_cast(total); + if (read_range) { + storage_size = read_range->restore_size; + if (storage_size < 0 || + read_range->source_offset > uint64_t(storage_size) || + total > uint64_t(storage_size) - read_range->source_offset) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + } + sizes.emplace_back(storage_size); + } + + const bool local_batch = + can_use_pinned_restore_arena(target_rpc_service_addr, objects); + std::optional local_owner; + auto response = + [&]() -> tl::expected { + if (!local_batch) { + return client_requester_->batch_get_offload_object( + target_rpc_service_addr, storage_keys, sizes); + } + auto result = file_storage_->BatchGetLocal(storage_keys, sizes); + if (!result) return tl::make_unexpected(result.error()); + local_owner.emplace(std::move(result.value())); + return BatchGetOffloadObjectResponse(0, + std::move(local_owner->pointers), + client_->GetSegmentEndpoint(), 0); + }(); + if (!response) { LOG(ERROR) << "Batch get offload object failed with error: " - << batchGetResp.error(); - return tl::make_unexpected(batchGetResp.error()); + << response.error(); + return tl::make_unexpected(response.error()); } - if (batchGetResp->pointers.size() != keys.size()) { + + const auto release_buffer = [&]() { + if (!local_batch) { + client_requester_->release_offload_buffer(target_rpc_service_addr, + response->batch_id); + } + }; + struct ReleaseGuard { + const decltype(release_buffer) &release; + ~ReleaseGuard() { release(); } + } release_guard{release_buffer}; + if (response->pointers.size() != keys.size()) { LOG(ERROR) << "Pointer count mismatch from owner: expected=" - << keys.size() << ", got=" << batchGetResp->pointers.size(); + << keys.size() << ", got=" << response->pointers.size(); return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } - auto result = - client_->BatchGetOffloadObject(batchGetResp->transfer_engine_addr, keys, - batchGetResp->pointers, objects); + if (read_range) { + if (response->pointers[0] > + std::numeric_limits::max() - read_range->source_offset) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + response->pointers[0] += read_range->source_offset; + } + auto result = client_->BatchGetOffloadObject( + response->transfer_engine_addr, keys, response->pointers, objects, + local_batch ? OffloadBufferAccess::kLocalAddress + : OffloadBufferAccess::kTransferEngine); auto end_time = std::chrono::steady_clock::now(); auto elapsed_time = static_cast( std::chrono::duration_cast(end_time - @@ -5961,20 +6051,15 @@ RealClient::batch_get_into_offload_object_internal( << elapsed_time << "ms, with target_rpc_service_addr: " << target_rpc_service_addr << ", key size: " << objects.size() - << ", batch_id: " << batchGetResp->batch_id - << ", gc ttl: " << batchGetResp->gc_ttl_ms << "ms."; - - // Release buffer immediately after transfer completion (fire-and-forget) - // This allows early buffer reclamation instead of waiting for GC lease - client_requester_->release_offload_buffer(target_rpc_service_addr, - batchGetResp->batch_id); + << ", batch_id: " << response->batch_id + << ", gc ttl: " << response->gc_ttl_ms << "ms."; if (!result) { LOG(ERROR) << "Batch get into offload object failed with error: " << result.error(); return result; } - if (elapsed_time >= batchGetResp->gc_ttl_ms) { + if (!local_batch && elapsed_time >= response->gc_ttl_ms) { return tl::make_unexpected(ErrorCode::OBJECT_HAS_LEASE); } return {}; diff --git a/mooncake-store/src/transfer_task.cpp b/mooncake-store/src/transfer_task.cpp index ca7faef384..134472b256 100644 --- a/mooncake-store/src/transfer_task.cpp +++ b/mooncake-store/src/transfer_task.cpp @@ -8,6 +8,7 @@ #include #include #include +#include #include #include #include @@ -1086,48 +1087,73 @@ std::optional TransferSubmitter::submit_batch_get_offload_object( const std::string& transfer_engine_addr, const std::vector& keys, const std::vector& pointers, - const std::unordered_map>& batched_slices) { - std::optional future; - std::vector requests; - // Open the segment once — all keys share the same transfer_engine_addr. - SegmentHandle seg = engine_.openSegment(transfer_engine_addr); - if (seg == static_cast(ERR_INVALID_ARGUMENT)) { - LOG(ERROR) << "Failed to open segment " << transfer_engine_addr; - // nullopt = failure (caller checks !future). The function returns - // std::optional so tl::unexpected is not available here. + const std::unordered_map>& batched_slices, + OffloadBufferAccess buffer_access) { + if (keys.size() != pointers.size()) { + LOG(ERROR) << "Mismatched offload transfer argument counts"; return std::nullopt; } + + const bool use_local_memcpy = + buffer_access == OffloadBufferAccess::kLocalAddress; + if (use_local_memcpy && !canUseLocalMemcpy(transfer_engine_addr)) { + LOG(ERROR) << "Offload source is not locally addressable: " + << transfer_engine_addr; + return std::nullopt; + } + + std::vector requests; + std::vector operations; + constexpr uint64_t kMaxAddress = std::numeric_limits::max(); + SegmentHandle seg = 0; + if (!use_local_memcpy) { + // Open once: all keys share the same transfer endpoint. + seg = engine_.openSegment(transfer_engine_addr); + if (seg == static_cast(ERR_INVALID_ARGUMENT)) { + LOG(ERROR) << "Failed to open segment " << transfer_engine_addr; + return std::nullopt; + } + } + for (size_t i = 0; i < keys.size(); ++i) { const auto& key = keys[i]; - const uint64_t pointer = pointers[i]; auto it = batched_slices.find(key); if (it == batched_slices.end()) { LOG(ERROR) << "Key not found in batched_slices: " << key; - return std::nullopt; // fail closed + return std::nullopt; } - // Emit one TransferRequest per slice: the on-disk blob is read - // sequentially while slices may point to non-contiguous GPU memory. uint64_t offset = 0; for (const auto& slice : it->second) { - TransferRequest request; - request.opcode = TransferRequest::READ; - request.source = static_cast(slice.ptr); - request.target_id = seg; - request.target_offset = pointer + offset; - request.length = slice.size; - requests.emplace_back(request); + if (slice.size == 0) continue; + if (!slice.ptr || pointers[i] > kMaxAddress - offset || + slice.size > kMaxAddress - pointers[i] - offset) { + LOG(ERROR) << "Invalid offload transfer range for key: " << key; + return std::nullopt; + } + if (use_local_memcpy) { + operations.emplace_back( + slice.ptr, + reinterpret_cast(pointers[i] + offset), + slice.size); + } else { + requests.emplace_back(TransferRequest{ + .opcode = TransferRequest::READ, + .source = static_cast(slice.ptr), + .target_id = seg, + .target_offset = pointers[i] + offset, + .length = slice.size, + }); + } offset += slice.size; } } - return submitTransfer(requests); + return use_local_memcpy ? submitMemcpyOperations(std::move(operations)) + : submitTransfer(requests); } std::optional TransferSubmitter::submitMemcpyOperation( const AllocatedBuffer::Descriptor& handle, const std::vector& slices, const TransferRequest::OpCode op_code, uint64_t src_offset) { - auto state = std::make_shared(); - - // Create memcpy operations std::vector operations; operations.reserve(slices.size()); uint64_t base_address = static_cast(handle.buffer_address_); @@ -1155,12 +1181,18 @@ std::optional TransferSubmitter::submitMemcpyOperation( operations.emplace_back(dest, src, slice.size); } - // Submit memcpy operations to worker pool for async execution + return submitMemcpyOperations(std::move(operations)); +} + +std::optional TransferSubmitter::submitMemcpyOperations( + std::vector operations) { + auto state = std::make_shared(); + const size_t operation_count = operations.size(); MemcpyTask task(std::move(operations), state); memcpy_pool_->submitTask(std::move(task)); - VLOG(1) << "Memcpy transfer submitted to worker pool with " << slices.size() - << " operations"; + VLOG(1) << "Memcpy transfer submitted to worker pool with " + << operation_count << " operations"; return TransferFuture(state); } @@ -1352,50 +1384,17 @@ std::optional TransferSubmitter::submitFileReadOperation( TransferStrategy TransferSubmitter::selectStrategy( const AllocatedBuffer::Descriptor& handle, - const std::vector& slices) const { - // Check if memcpy operations are enabled via environment variable - if (!memcpy_enabled_) { - VLOG(2) << "Memcpy operations disabled via MC_STORE_MEMCPY environment " - "variable"; - return TransferStrategy::TRANSFER_ENGINE; - } - - // Check conditions for local memcpy optimization - if (isLocalTransfer(handle)) { - return TransferStrategy::LOCAL_MEMCPY; - } - - return TransferStrategy::TRANSFER_ENGINE; + const std::vector& /* slices */) const { + return canUseLocalMemcpy(handle.transport_endpoint_) + ? TransferStrategy::LOCAL_MEMCPY + : TransferStrategy::TRANSFER_ENGINE; } -namespace { -// Helper function to extract IP address from endpoint string (ip:port format). -// Supports both IPv4 (ip:port) and IPv6 ([ipv6]:port) formats. -std::string extractIpAddress(const std::string& endpoint) { - if (endpoint.empty()) { - return ""; - } - - // Handle IPv6 format: [ipv6]:port - if (endpoint[0] == '[') { - size_t closing_bracket = endpoint.find(']'); - if (closing_bracket == std::string::npos) { - LOG(WARNING) << "Invalid IPv6 endpoint format: " << endpoint; - return ""; - } - return endpoint.substr(1, closing_bracket - 1); - } - - // Handle IPv4 or hostname:port format. - size_t colon_pos = endpoint.rfind(':'); - if (colon_pos != std::string::npos) { - return endpoint.substr(0, colon_pos); - } - - // No colon found, return the whole string (might be just IP or hostname). - return endpoint; +bool TransferSubmitter::canUseLocalMemcpy(const std::string& endpoint) const { + return memcpy_enabled_ && + (isSameProcessEndpoint(endpoint, local_hostname_) || + isSameProcessEndpoint(endpoint, local_endpoint_)); } -} // namespace bool TransferSubmitter::isSameProcessEndpoint( const std::string& handle_endpoint, const std::string& local_endpoint) { @@ -1405,28 +1404,7 @@ bool TransferSubmitter::isSameProcessEndpoint( // memcpy on a peer process's address would segfault. Require the full // transport endpoint to match, which uniquely identifies the owning // process. - if (handle_endpoint.empty() || local_endpoint.empty()) { - return false; - } - if (handle_endpoint == local_endpoint) { - return true; - } - - const std::string handle_ip = extractIpAddress(handle_endpoint); - const std::string local_ip = extractIpAddress(local_endpoint); - if (!handle_ip.empty() && handle_ip == local_ip) { - VLOG(2) << "Disabling local memcpy for same-host endpoints with " - "different process endpoints: handle=" - << handle_endpoint << ", local=" << local_endpoint; - } - - return false; -} - -bool TransferSubmitter::isLocalTransfer( - const AllocatedBuffer::Descriptor& handle) const { - return isSameProcessEndpoint(handle.transport_endpoint_, local_hostname_) || - isSameProcessEndpoint(handle.transport_endpoint_, local_endpoint_); + return !handle_endpoint.empty() && handle_endpoint == local_endpoint; } bool TransferSubmitter::validateTransferParams( diff --git a/mooncake-store/tests/CMakeLists.txt b/mooncake-store/tests/CMakeLists.txt index 201ab04c87..72e18be31f 100644 --- a/mooncake-store/tests/CMakeLists.txt +++ b/mooncake-store/tests/CMakeLists.txt @@ -106,17 +106,23 @@ add_store_test(master_admin_server_test master_admin_server_test.cpp) add_store_test(posix_file_test posix_file_test.cpp) add_store_test(thread_pool_test thread_pool_test.cpp) add_store_test(transfer_task_test transfer_task_test.cpp) +find_package(CUDAToolkit QUIET) +if(CUDAToolkit_FOUND) + target_compile_definitions(transfer_task_test PRIVATE MOONCAKE_TEST_CUDA_H2D) + target_link_libraries(transfer_task_test PRIVATE CUDA::cudart) +endif() if(USE_CUDA) find_package(CUDAToolkit REQUIRED) target_compile_definitions(transfer_task_test PRIVATE USE_CUDA) - target_link_libraries(transfer_task_test PRIVATE CUDA::cudart) if(USE_TENT) add_test( NAME transfer_scatter_tent_test - COMMAND transfer_task_test - --gtest_filter=TransferTaskTest.TransferScatterWritesGpuDestinationDirectly) - set_tests_properties(transfer_scatter_tent_test - PROPERTIES ENVIRONMENT "MC_USE_TENT=1") + COMMAND + transfer_task_test + --gtest_filter=TransferTaskTest.TransferScatterWritesGpuDestinationDirectly + ) + set_tests_properties(transfer_scatter_tent_test PROPERTIES ENVIRONMENT + "MC_USE_TENT=1") endif() endif() add_store_test(tenant_quota_test tenant_quota_test.cpp) @@ -128,6 +134,10 @@ add_store_test(client_buffer_test client_buffer_test.cpp) add_store_test(client_local_hot_cache_test client_local_hot_cache_test.cpp) add_store_test(client_tcp_local_memcpy_test client_tcp_local_memcpy_test.cpp) add_store_test(pybind_client_test pybind_client_test.cpp) +if(CUDAToolkit_FOUND) + target_compile_definitions(pybind_client_test PRIVATE MOONCAKE_TEST_CUDA_H2D) + target_link_libraries(pybind_client_test PRIVATE CUDA::cudart) +endif() add_store_test(ipv6_client_test ipv6_client_test.cpp) add_store_test(host_port_fix_test host_port_fix_test.cpp) add_store_test(client_metrics_test client_metrics_test.cpp) diff --git a/mooncake-store/tests/file_storage_test.cpp b/mooncake-store/tests/file_storage_test.cpp index ef29ccae8b..15f38c9d14 100644 --- a/mooncake-store/tests/file_storage_test.cpp +++ b/mooncake-store/tests/file_storage_test.cpp @@ -30,6 +30,7 @@ class FileStorageTest : public ::testing::Test { FLAGS_logtostderr = true; UnsetEnv("MOONCAKE_OFFLOAD_FILE_STORAGE_PATH"); UnsetEnv("MOONCAKE_OFFLOAD_LOCAL_BUFFER_SIZE_BYTES"); + UnsetEnv("MC_STORE_PINNED_RESTORE_ARENA_SIZE_BYTES"); UnsetEnv("MOONCAKE_OFFLOAD_SCANMETA_ITERATOR_KEYS_LIMIT"); UnsetEnv("MOONCAKE_SCANMETA_ITERATOR_KEYS_LIMIT"); UnsetEnv("MOONCAKE_OFFLOAD_BUCKET_KEYS_LIMIT"); @@ -67,7 +68,14 @@ class FileStorageTest : public ::testing::Test { FileStorageAllocateBatch(FileStorage& fileStorage, const std::vector& keys, const std::vector& sizes) { - return fileStorage.AllocateBatch(keys, sizes); + return fileStorage.AllocateBatch(keys, sizes, + *fileStorage.client_buffer_allocator_); + } + + void SetPinnedRestoreArena(FileStorage& fileStorage, void* address, + size_t size) { + fileStorage.pinned_restore_arena_allocator_ = + ClientBufferAllocator::create(address, size); } tl::expected FileStorageBatchLoad( @@ -216,27 +224,45 @@ TEST_F(FileStorageTest, IsEnableOffloading) { !enable_offloading_result3.value()); } -TEST_F(FileStorageTest, BatchLoad) { +TEST_F(FileStorageTest, BatchGetUsesPinnedArenaAndFallsBackWhenFull) { std::vector keys; std::vector sizes; std::unordered_map batch_data; + auto file_storage_config = FileStorageConfig::FromEnvironment(); file_storage_config.storage_filepath = data_path; + file_storage_config.local_buffer_size = 128 * 1024 * 1024; + constexpr size_t kArenaSize = 2 * 1024 * 1024; + std::vector restore_arena(kArenaSize + 4096); FileStorage fileStorage(file_storage_config, nullptr, "localhost:9003"); + const auto arena_begin = + (reinterpret_cast(restore_arena.data()) + 4095) & ~4095ULL; + SetPinnedRestoreArena(fileStorage, reinterpret_cast(arena_begin), + kArenaSize); ASSERT_TRUE(FileStorageBatchOffload(fileStorage, keys, sizes, batch_data)); - std::unordered_map batch_slice; - std::vector buff; - auto allocate_res = FileStorageAllocateBatch(fileStorage, keys, sizes); - ASSERT_TRUE(allocate_res); + auto pinned_result = fileStorage.BatchGetLocal(keys, sizes); + ASSERT_TRUE(pinned_result); + ASSERT_EQ(pinned_result->pointers.size(), keys.size()); + for (size_t i = 0; i < keys.size(); ++i) { + EXPECT_GE(pinned_result->pointers[i], arena_begin); + EXPECT_LT(pinned_result->pointers[i], arena_begin + kArenaSize); + EXPECT_EQ( + std::string(reinterpret_cast(pinned_result->pointers[i]), + sizes[i]), + batch_data.at(keys[i])); + } - ASSERT_TRUE( - FileStorageBatchLoad(fileStorage, allocate_res.value()->slices)); - for (auto& slice_it : batch_slice) { - std::string data(static_cast(slice_it.second.ptr), - slice_it.second.size); - LOG(INFO) << "key: " << slice_it.first; - ASSERT_EQ(data, batch_data.at(slice_it.first)); + auto fallback_result = fileStorage.BatchGetLocal(keys, sizes); + ASSERT_TRUE(fallback_result); + ASSERT_EQ(fallback_result->pointers.size(), keys.size()); + for (size_t i = 0; i < keys.size(); ++i) { + EXPECT_TRUE(fallback_result->pointers[i] < arena_begin || + fallback_result->pointers[i] >= arena_begin + kArenaSize); + EXPECT_EQ( + std::string(reinterpret_cast(fallback_result->pointers[i]), + sizes[i]), + batch_data.at(keys[i])); } } @@ -357,6 +383,7 @@ TEST_F(FileStorageTest, DefaultValuesWhenNoEnvSet) { EXPECT_EQ(config.total_size_limit, 2ULL * 1024 * 1024 * 1024 * 1024); EXPECT_EQ(config.heartbeat_interval_seconds, 10u); EXPECT_TRUE(config.enable_disk_watermark_eviction); + EXPECT_EQ(config.pinned_restore_arena_size, 0); EXPECT_DOUBLE_EQ(config.disk_eviction_high_watermark_ratio, 0.90); EXPECT_DOUBLE_EQ(config.disk_eviction_low_watermark_ratio, 0.80); } @@ -370,6 +397,7 @@ TEST_F(FileStorageTest, ReadStringFromEnv) { TEST_F(FileStorageTest, ReadInt64FromEnv) { SetEnv("MOONCAKE_OFFLOAD_LOCAL_BUFFER_SIZE_BYTES", "2147483648"); // 2GB + SetEnv("MC_STORE_PINNED_RESTORE_ARENA_SIZE_BYTES", "67108864"); SetEnv("MOONCAKE_OFFLOAD_BUCKET_KEYS_LIMIT", "1000"); SetEnv("MOONCAKE_OFFLOAD_TOTAL_KEYS_LIMIT", "5000000"); @@ -377,6 +405,7 @@ TEST_F(FileStorageTest, ReadInt64FromEnv) { auto bucket_backend_config = BucketBackendConfig::FromEnvironment(); EXPECT_EQ(config.local_buffer_size, 2147483648); + EXPECT_EQ(config.pinned_restore_arena_size, 64 * 1024 * 1024); EXPECT_EQ(bucket_backend_config.bucket_keys_limit, 1000); EXPECT_EQ(config.total_keys_limit, 5000000); } @@ -623,6 +652,10 @@ TEST_F(FileStorageTest, ValidateFailsOnInvalidLimits) { EXPECT_FALSE(config.Validate()); config.total_size_limit = 1; + config.pinned_restore_arena_size = -1; + EXPECT_FALSE(config.Validate()); + + config.pinned_restore_arena_size = 0; config.heartbeat_interval_seconds = 0; EXPECT_FALSE(config.Validate()); diff --git a/mooncake-store/tests/pybind_client_test.cpp b/mooncake-store/tests/pybind_client_test.cpp index c08a48a8e7..b9c7057097 100644 --- a/mooncake-store/tests/pybind_client_test.cpp +++ b/mooncake-store/tests/pybind_client_test.cpp @@ -6,14 +6,20 @@ #include #include #include +#include #include #include #include +#include #include #include #include #include +#ifdef MOONCAKE_TEST_CUDA_H2D +#include +#endif + #include "real_client.h" #include "test_server_helpers.h" @@ -36,6 +42,26 @@ class GLogMuter { int original_log_level_; }; +class ScopedEnvVar { + public: + ScopedEnvVar(const char* name, const char* value) : name_(name) { + if (const char* previous = getenv(name)) previous_ = previous; + setenv(name, value, 1); + } + + ~ScopedEnvVar() { + if (previous_) { + setenv(name_.c_str(), previous_->c_str(), 1); + } else { + unsetenv(name_.c_str()); + } + } + + private: + std::string name_; + std::optional previous_; +}; + class RealClientTest : public ::testing::Test { protected: static void SetUpTestSuite() { @@ -60,9 +86,11 @@ class RealClientTest : public ::testing::Test { void TearDown() override { if (py_client_) { py_client_->tearDownAll(); + if (!ssd_path_.empty()) py_client_.reset(); } master_.Stop(); + if (!ssd_path_.empty()) std::filesystem::remove_all(ssd_path_); } std::shared_ptr py_client_; @@ -70,6 +98,7 @@ class RealClientTest : public ::testing::Test { // In-proc master for tests mooncake::testing::InProcMaster master_; std::string master_address_; + std::string ssd_path_; void StartMasterAndSetupClient() { ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) @@ -117,6 +146,113 @@ class RealClientTest : public ::testing::Test { } }; +#ifdef MOONCAKE_TEST_CUDA_H2D +TEST_F(RealClientTest, PinnedSsdRestoreReadsNonTailRangeIntoGpu) { + int device_count = 0; + if (cudaGetDeviceCount(&device_count) != cudaSuccess || device_count == 0) { + GTEST_SKIP() << "CUDA device is unavailable"; + } + + ScopedEnvVar local_memcpy("MC_STORE_MEMCPY", "1"); + ScopedEnvVar heartbeat("MOONCAKE_OFFLOAD_HEARTBEAT_INTERVAL_SECONDS", "1"); + ScopedEnvVar pinned_restore_arena( + "MC_STORE_PINNED_RESTORE_ARENA_SIZE_BYTES", "1048576"); + ScopedEnvVar storage_backend("MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR", + "bucket_storage_backend"); + ScopedEnvVar bucket_keys("MOONCAKE_OFFLOAD_BUCKET_KEYS_LIMIT", "1"); + + char path[] = "/tmp/mooncake_ssd_pinned_restore_XXXXXX"; + const char* created = mkdtemp(path); + ASSERT_NE(created, nullptr); + ssd_path_ = created; + + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder() + .set_enable_offload(true) + .set_default_kv_lease_ttl(10) + .build())); + master_address_ = master_.master_address(); + ASSERT_EQ( + py_client_->setup_real("localhost:17813", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, "tcp", "", + master_address_, nullptr, "", true, ssd_path_), + 0); + + constexpr size_t kObjectSize = 64 * 1024; + constexpr size_t kSourceOffset = 8 * 1024; + constexpr size_t kRangeSize = 16 * 1024; + static_assert(kSourceOffset + kRangeSize < kObjectSize); + std::vector source(kObjectSize); + for (size_t i = 0; i < source.size(); ++i) { + source[i] = static_cast(i % 251); + } + const std::string key = "pinned_ssd_non_tail_range"; + ASSERT_EQ(py_client_->put(key, source), 0); + + bool disk_ready = false; + const auto offload_deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(10); + while (std::chrono::steady_clock::now() < offload_deadline && !disk_ready) { + for (const auto& replica : py_client_->get_replica_desc(key)) { + if (replica.is_local_disk_replica()) { + disk_ready = true; + } + } + if (!disk_ready) + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + ASSERT_TRUE(disk_ready); + + std::this_thread::sleep_for(std::chrono::milliseconds(20)); + const auto clear_deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(2); + bool memory_cleared = false; + while (std::chrono::steady_clock::now() < clear_deadline) { + if (py_client_->batch_replica_clear({key}, "localhost:17813").size() == + 1) { + memory_cleared = true; + break; + } + std::this_thread::sleep_for(std::chrono::milliseconds(20)); + } + ASSERT_TRUE(memory_cleared); + + auto replicas = py_client_->get_replica_desc(key); + ASSERT_EQ(replicas.size(), 1); + ASSERT_TRUE(replicas.front().is_local_disk_replica()); + + void* gpu_destination = nullptr; + ASSERT_EQ(cudaMalloc(&gpu_destination, kRangeSize), cudaSuccess); + bool registered = false; + auto cleanup = [this, ®istered](void* ptr) { + if (registered) { + EXPECT_EQ(py_client_->unregister_buffer(ptr), 0); + } + EXPECT_EQ(cudaFree(ptr), cudaSuccess); + }; + std::unique_ptr gpu_owner(gpu_destination, + cleanup); + ASSERT_EQ(py_client_->register_buffer(gpu_destination, kRangeSize), 0); + registered = true; + + const auto reads_before = py_client_->get_offload_rpc_read_count(); + auto results = + py_client_->get_into_ranges({gpu_destination}, {{key}}, {{{0}}}, + {{{kSourceOffset}}}, {{{kRangeSize}}}); + ASSERT_EQ(results.size(), 1); + ASSERT_EQ(results[0].size(), 1); + ASSERT_EQ(results[0][0].size(), 1); + EXPECT_EQ(results[0][0][0], static_cast(kRangeSize)); + EXPECT_EQ(py_client_->get_offload_rpc_read_count(), reads_before + 1); + + std::vector actual(kRangeSize); + ASSERT_EQ(cudaMemcpy(actual.data(), gpu_destination, kRangeSize, + cudaMemcpyDeviceToHost), + cudaSuccess); + EXPECT_TRUE(std::equal(actual.begin(), actual.end(), + source.begin() + kSourceOffset)); +} +#endif + TEST_F(RealClientTest, AllocateAndMountSegmentAlignsAndUnmounts) { StartMasterAndSetupClient(); diff --git a/mooncake-store/tests/transfer_task_test.cpp b/mooncake-store/tests/transfer_task_test.cpp index 8753757ca0..7ec18f635f 100644 --- a/mooncake-store/tests/transfer_task_test.cpp +++ b/mooncake-store/tests/transfer_task_test.cpp @@ -6,30 +6,32 @@ #include #include +#include #include #include #include #include #include "types.h" -#ifdef USE_CUDA +#include "pinned_buffer_pool.h" +#if defined(USE_CUDA) || defined(MOONCAKE_TEST_CUDA_H2D) #include #endif namespace mooncake { // Test fixture for TransferTask tests -// TODO: Currently, this test does not cover TransferSubmitter and -// TransferEngine integration. Will add more tests in the future. class TransferTaskTest : public ::testing::Test { protected: void SetUp() override { // Initialize glog for logging google::InitGoogleLogging("TransferTaskTest"); FLAGS_logtostderr = 1; // Output logs to stderr + unsetenv("MC_STORE_MEMCPY"); } void TearDown() override { + unsetenv("MC_STORE_MEMCPY"); // Cleanup glog google::ShutdownGoogleLogging(); } @@ -281,11 +283,7 @@ TEST_F(TransferTaskTest, TransferScatterHandlesFragmentedGpuBuffers) { } #endif -// Test the locality decision used by TransferSubmitter::isLocalTransfer. -// Same-host different-process pairs share an IP but have distinct ports; -// they must NOT be treated as locally addressable, otherwise memcpy in the -// caller process would dereference a virtual address belonging to a peer -// process and segfault. +// Same-host endpoints from different processes are not locally addressable. TEST_F(TransferTaskTest, IsSameProcessEndpoint) { // Empty inputs -> not same-process (cannot prove locality). EXPECT_FALSE(TransferSubmitter::isSameProcessEndpoint("", "")); @@ -312,6 +310,88 @@ TEST_F(TransferTaskTest, IsSameProcessEndpoint) { EXPECT_FALSE(TransferSubmitter::isSameProcessEndpoint("host-a", "host-b")); } +TEST_F(TransferTaskTest, BatchGetOffloadObjectHonorsLocalMemcpySetting) { + setenv("MC_STORE_MEMCPY", "1", 1); + TransferEngine engine(false); + ASSERT_EQ(engine.init("P2PHANDSHAKE", "localhost:17933"), 0); + const std::string endpoint = engine.getLocalIpAndPort(); + + std::vector source(512, 'A'); + std::vector destination(512, 0); + const std::vector keys{"key"}; + const std::vector pointers{ + reinterpret_cast(source.data())}; + const std::unordered_map> slices{ + {"key", {{nullptr, 0}, {destination.data(), destination.size()}}}}; + + { + std::shared_ptr backend; + TransferSubmitter submitter(engine, backend, endpoint); + auto future = submitter.submit_batch_get_offload_object( + endpoint, keys, pointers, slices, + OffloadBufferAccess::kLocalAddress); + ASSERT_TRUE(future); + EXPECT_EQ(future->strategy(), TransferStrategy::LOCAL_MEMCPY); + EXPECT_EQ(future->get(), ErrorCode::OK); + EXPECT_FALSE(submitter.submit_batch_get_offload_object( + endpoint, keys, {std::numeric_limits::max() - 7}, + {{"key", {{destination.data(), 16}}}}, + OffloadBufferAccess::kLocalAddress)); + } + EXPECT_EQ(destination, source); + + setenv("MC_STORE_MEMCPY", "0", 1); + std::shared_ptr backend; + TransferSubmitter submitter(engine, backend, endpoint); + EXPECT_FALSE(submitter.submit_batch_get_offload_object( + endpoint, keys, pointers, slices, OffloadBufferAccess::kLocalAddress)); + EXPECT_EQ(engine.freeEngine(), 0); +} + +#ifdef MOONCAKE_TEST_CUDA_H2D +TEST_F(TransferTaskTest, BatchGetOffloadObjectCopiesPinnedHostToGpu) { + int device_count = 0; + if (cudaGetDeviceCount(&device_count) != cudaSuccess || device_count == 0) { + GTEST_SKIP() << "CUDA device is unavailable"; + } + + setenv("MC_STORE_MEMCPY", "1", 1); + constexpr size_t kSourceOffset = 128; + constexpr size_t kSize = 4096; + auto pinned_buffer = + PinnedBufferPool::AllocatePinned(kSourceOffset + kSize); + ASSERT_NE(pinned_buffer.pinned_host.addr, nullptr); + void* pinned_source = pinned_buffer.data; + void* gpu_destination = nullptr; + ASSERT_EQ(cudaMalloc(&gpu_destination, kSize), cudaSuccess); + std::memset(static_cast(pinned_source) + kSourceOffset, 0x5a, kSize); + + TransferEngine engine(false); + ASSERT_EQ(engine.init("P2PHANDSHAKE", "localhost:17934"), 0); + const std::string endpoint = engine.getLocalIpAndPort(); + { + std::shared_ptr backend; + TransferSubmitter submitter(engine, backend, endpoint); + auto future = submitter.submit_batch_get_offload_object( + endpoint, {"gpu"}, + {reinterpret_cast(static_cast(pinned_source) + + kSourceOffset)}, + {{"gpu", {{gpu_destination, kSize}}}}, + OffloadBufferAccess::kLocalAddress); + ASSERT_TRUE(future); + EXPECT_EQ(future->get(), ErrorCode::OK); + } + + std::vector actual(kSize); + EXPECT_EQ(cudaMemcpy(actual.data(), gpu_destination, kSize, + cudaMemcpyDeviceToHost), + cudaSuccess); + EXPECT_EQ(actual, std::vector(kSize, 0x5a)); + EXPECT_EQ(engine.freeEngine(), 0); + EXPECT_EQ(cudaFree(gpu_destination), cudaSuccess); +} +#endif + // Test TransferStrategy enum and stream operator TEST_F(TransferTaskTest, TransferStrategyEnum) { // Test enum values From d82b9bc19e4d86a832de13fe8749faeafb1595a9 Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Mon, 10 Aug 2026 10:54:49 +0800 Subject: [PATCH 009/483] [Store] Fix flaky standby catch-up test (#3273) Co-authored-by: Yuchen Kou --- .../ha/standby/hot_standby_service_test.cpp | 54 +++++++++++++++---- 1 file changed, 43 insertions(+), 11 deletions(-) diff --git a/mooncake-store/tests/ha/standby/hot_standby_service_test.cpp b/mooncake-store/tests/ha/standby/hot_standby_service_test.cpp index 317a48d799..2bb20ac045 100644 --- a/mooncake-store/tests/ha/standby/hot_standby_service_test.cpp +++ b/mooncake-store/tests/ha/standby/hot_standby_service_test.cpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include #include @@ -558,6 +559,7 @@ std::string MakeValidPayload(uint64_t client_id_first = 1, class FakeHaKvBackend : public HaKvBackend { public: ErrorCode Get(std::string_view key, std::string& value) override { + std::lock_guard lock(mutex_); if (next_get_error_ != ErrorCode::OK) { auto error = next_get_error_; next_get_error_ = ErrorCode::OK; @@ -574,11 +576,13 @@ class FakeHaKvBackend : public HaKvBackend { return ErrorCode::OK; } ErrorCode Put(std::string_view key, std::string_view value) override { + std::lock_guard lock(mutex_); values_[std::string(key)] = std::string(value); return ErrorCode::OK; } ErrorCode Range(std::string_view begin_key, std::string_view end_key, size_t limit, std::vector& kvs) override { + std::lock_guard lock(mutex_); if (range_error_ != ErrorCode::OK) { return range_error_; } @@ -592,12 +596,39 @@ class FakeHaKvBackend : public HaKvBackend { return ErrorCode::OK; } bool SupportsTxn() const override { return true; } - ErrorCode Txn(const KvTxn&) override { return ErrorCode::OK; } - void SetGetError(ErrorCode err) { get_error_ = err; } - void SetRangeError(ErrorCode err) { range_error_ = err; } - void FailNextGet(ErrorCode err) { next_get_error_ = err; } + ErrorCode Txn(const KvTxn& txn) override { + std::lock_guard lock(mutex_); + for (const auto& compare : txn.compares) { + auto it = values_.find(compare.key); + if (compare.kind == KvCompareKind::kKeyNotExists) { + if (it != values_.end()) { + return ErrorCode::ETCD_TRANSACTION_FAIL; + } + } else if (it == values_.end() || + it->second != compare.expected_value) { + return ErrorCode::ETCD_TRANSACTION_FAIL; + } + } + for (const auto& put : txn.puts) { + values_[put.key] = put.value; + } + return ErrorCode::OK; + } + void SetGetError(ErrorCode err) { + std::lock_guard lock(mutex_); + get_error_ = err; + } + void SetRangeError(ErrorCode err) { + std::lock_guard lock(mutex_); + range_error_ = err; + } + void FailNextGet(ErrorCode err) { + std::lock_guard lock(mutex_); + next_get_error_ = err; + } private: + mutable std::mutex mutex_; std::map values_; ErrorCode get_error_{ErrorCode::OK}; ErrorCode range_error_{ErrorCode::OK}; @@ -870,16 +901,17 @@ TEST_F(PromotionCatchUpTest, MissingDurablePrefixRejectsNonzeroSequence) { } TEST_F(PromotionCatchUpTest, CatchesUpPrefixThatAppearsBeforePromotion) { + const DurablePrefix initial_prefix{.batch_id = 0, .last_seq = 0}; + ASSERT_EQ(ErrorCode::OK, + batch_backend_->Put(BuildDurablePrefixKey(cluster_id_), + EncodeDurablePrefix(initial_prefix))); ASSERT_EQ(ErrorCode::OK, service_->Start("", oplog_endpoints_, cluster_id_)); ASSERT_EQ(StandbyState::WATCHING, service_->GetState()); - ASSERT_EQ(ErrorCode::OK, - batch_backend_->Put( - BuildDurablePrefixKey(cluster_id_), - EncodeDurablePrefix({.batch_id = 1, .last_seq = 1}))); - ASSERT_EQ(ErrorCode::OK, - batch_backend_->Put(BuildBatchRecordKey(cluster_id_, 1), - EncodeOpLogBatchRecord(MakeBatch(1, 1, 1)))); + + OpLogBatchStorage storage(cluster_id_, *batch_backend_); + ASSERT_EQ(ErrorCode::OK, storage.WriteBatchAndAdvancePrefix( + MakeBatch(1, 1, 1), initial_prefix)); StandbySnapshot out; ASSERT_EQ(ErrorCode::OK, service_->PromoteAndExportSnapshot(out)); From cab1c42676239b6a3677127a914090d529e1c51c Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Mon, 10 Aug 2026 10:59:20 +0800 Subject: [PATCH 010/483] [TE] Add IBGDA per-QP GPU-VA control path (#3334) --- .../include/gpu_vendor/musa.h | 14 ++ .../include/transport/device/ibgda/mlx5gda.h | 41 ++-- .../include/transport/device/ibgda_device.cuh | 29 ++- .../device/ibgda_device_transport.cpp | 227 ++++++++++++++++-- .../src/transport/device/mlx5gda.cpp | 189 ++++++++------- 5 files changed, 370 insertions(+), 130 deletions(-) diff --git a/mooncake-transfer-engine/include/gpu_vendor/musa.h b/mooncake-transfer-engine/include/gpu_vendor/musa.h index 4fd47def0b..7894abf21a 100644 --- a/mooncake-transfer-engine/include/gpu_vendor/musa.h +++ b/mooncake-transfer-engine/include/gpu_vendor/musa.h @@ -2,6 +2,7 @@ #include #include +#include #include const static std::string GPU_PREFIX = "musa:"; @@ -60,6 +61,7 @@ const static std::string GPU_PREFIX = "musa:"; #define CUDA_ERROR_NOT_SUPPORTED MUSA_ERROR_NOT_SUPPORTED #define cudaDeviceCanAccessPeer musaDeviceCanAccessPeer #define cudaDeviceEnablePeerAccess musaDeviceEnablePeerAccess +#define cudaDeviceGetStreamPriorityRange musaDeviceGetStreamPriorityRange #define cudaDeviceGetPCIBusId musaDeviceGetPCIBusId #define cudaErrorPeerAccessAlreadyEnabled musaErrorPeerAccessAlreadyEnabled #define cudaError_t musaError_t @@ -67,10 +69,14 @@ const static std::string GPU_PREFIX = "musa:"; #define cudaFreeHost musaFreeHost #define cudaGetDevice musaGetDevice #define cudaGetDeviceCount musaGetDeviceCount +#define cudaGetErrorName musaGetErrorName #define cudaGetErrorString musaGetErrorString #define cudaGetLastError musaGetLastError #define cudaHostAlloc musaHostAlloc +#define cudaHostAllocDefault musaHostAllocDefault #define cudaHostAllocMapped musaHostAllocMapped +#define cudaHostAllocPortable musaHostAllocPortable +#define cudaHostAllocWriteCombined musaHostAllocWriteCombined #define cudaHostRegister musaHostRegister #define cudaHostRegisterPortable musaHostRegisterPortable #define cudaHostUnregister musaHostUnregister @@ -97,6 +103,7 @@ const static std::string GPU_PREFIX = "musa:"; #define cudaSetDevice musaSetDevice #define cudaStreamCreate musaStreamCreate #define cudaStreamCreateWithFlags musaStreamCreateWithFlags +#define cudaStreamCreateWithPriority musaStreamCreateWithPriority #define cudaStreamNonBlocking musaStreamNonBlocking #define cudaStreamDestroy musaStreamDestroy #define cudaStreamPerThread musaStreamPerThread @@ -122,6 +129,10 @@ const static std::string GPU_PREFIX = "musa:"; #define cudaGetDeviceProperties musaGetDeviceProperties #define cudaMemcpyDeviceToDevice musaMemcpyDeviceToDevice #define cudaDevAttrClockRate musaDevAttrClockRate +#define cudaDevAttrMaxSharedMemoryPerBlockOptin \ + musaDevAttrMaxSharedMemoryPerBlockOptin +#define cudaEventCreate musaEventCreate +#define cudaEventElapsedTime musaEventElapsedTime #define cudaLaunchConfig_t musaLaunchConfig_t #define cudaLaunchAttribute musaLaunchAttribute #define cudaLaunchAttributeCooperative musaLaunchAttributeCooperative @@ -129,6 +140,9 @@ const static std::string GPU_PREFIX = "musa:"; #define CUDA_R_16BF MUSA_R_16BF #define CUDA_R_32F MUSA_R_32F +#define nv_bfloat16 __mt_bfloat16 +#define nv_bfloat162 __mt_bfloat162 + // IBGDA-specific mappings #define cuInit muInit #define cuDevicePrimaryCtxRetain muDevicePrimaryCtxRetain diff --git a/mooncake-transfer-engine/include/transport/device/ibgda/mlx5gda.h b/mooncake-transfer-engine/include/transport/device/ibgda/mlx5gda.h index 24154e6bf1..6fa7b5c887 100644 --- a/mooncake-transfer-engine/include/transport/device/ibgda/mlx5gda.h +++ b/mooncake-transfer-engine/include/transport/device/ibgda/mlx5gda.h @@ -3,12 +3,7 @@ #include #include -#ifdef USE_MUSA -#include -#define cudaStream_t musaStream_t -#else -#include -#endif +#include "cuda_alike.h" #include #include @@ -40,6 +35,13 @@ struct mlx5gda_control_region_allocator { void (*release)(void *context, struct mlx5gda_control_region *region); }; +struct mlx5gda_control_buffer { + void *addr; + void *dev_addr; + struct mlx5dv_devx_umem *umem; + struct memheap *heap; +}; + struct mlx5gda_rdma_write_wqe { struct mlx5_wqe_ctrl_seg ctrl; struct mlx5_wqe_raddr_seg raddr; @@ -58,6 +60,9 @@ struct mlx5gda_cq { struct mlx5dv_devx_uar *uar; // uar is allocated but not used uint32_t cqn; uint32_t cqe; + void *dev_base; + void *dev_cq_buf; + struct memheap *heap; size_t cq_offset; size_t dbr_offset; void *cq_buf; @@ -70,9 +75,10 @@ struct mlx5gda_cq { struct mlx5gda_cq *mlx5gda_create_cq( void *ctrl_buf, struct mlx5dv_devx_umem *ctrl_buf_umem, struct memheap *ctrl_buf_heap, struct ibv_pd *pd, int num_cqe, - cudaStream_t stream, - const struct mlx5gda_control_region_allocator *region_allocator); -void mlx5gda_destroy_cq(struct memheap *ctrl_buf_heap, struct mlx5gda_cq *cq); + cudaStream_t stream, void *cq_buf, void *cq_buf_dev, + struct mlx5dv_devx_umem *cq_buf_umem, struct memheap *cq_buf_heap, + const struct mlx5gda_control_region_allocator *cq_region_allocator); +void mlx5gda_destroy_cq(struct mlx5gda_cq *cq); static const size_t MLX5GDA_BF_SIZE = 256; @@ -80,6 +86,9 @@ struct mlx5gda_qp { struct mlx5dv_devx_obj *mqp; struct mlx5gda_cq *send_cq; struct mlx5dv_devx_uar *uar; + void *bf_host_base; + char *bf_device_addr; + bool bf_host_registration_owner; uint8_t port_num; struct ibv_port_attr port_attr; @@ -90,8 +99,11 @@ struct mlx5gda_qp { uint32_t num_wqebb; size_t wq_offset; size_t dbr_offset; + struct memheap *wq_heap; void *wq; void *dbr; + void *dev_wq; + void *dev_dbr; struct mlx5gda_control_region wq_region; struct mlx5gda_control_region dbr_region; struct mlx5gda_control_region_allocator region_allocator; @@ -121,11 +133,12 @@ void mlx5gda_reset_create_qp_failure(); mlx5gda_create_qp_failure mlx5gda_last_create_qp_failure(); struct mlx5gda_qp *mlx5gda_create_rc_qp( - struct mlx5dv_pd mpd, void *ctrl_buf, - struct mlx5dv_devx_umem *ctrl_buf_umem, struct memheap *ctrl_buf_heap, - struct ibv_pd *pd, int wqe, uint8_t port_num, cudaStream_t stream, - const struct mlx5gda_control_region_allocator *region_allocator); -void mlx5gda_destroy_qp(struct memheap *ctrl_buf_heap, struct mlx5gda_qp *qp); + struct mlx5dv_pd mpd, const struct mlx5gda_control_buffer *ctrl, + const struct mlx5gda_control_buffer *cq, struct ibv_pd *pd, int wqe, + uint8_t port_num, cudaStream_t stream, + const struct mlx5gda_control_region_allocator *cq_region_allocator, + const struct mlx5gda_control_region_allocator *qp_region_allocator); +void mlx5gda_destroy_qp(struct mlx5gda_qp *qp); int mlx5gda_modify_rc_qp_rst2init(struct mlx5gda_qp *qp, uint16_t pkey_index); int mlx5gda_modify_rc_qp_init2rtr(struct mlx5gda_qp *qp, diff --git a/mooncake-transfer-engine/include/transport/device/ibgda_device.cuh b/mooncake-transfer-engine/include/transport/device/ibgda_device.cuh index 2376f3b42e..df6964db2e 100644 --- a/mooncake-transfer-engine/include/transport/device/ibgda_device.cuh +++ b/mooncake-transfer-engine/include/transport/device/ibgda_device.cuh @@ -118,16 +118,41 @@ __device__ __forceinline__ void mc_ibgda_poll_cq(mlx5gda_qp_devctx* qp, __device__ __forceinline__ void mc_ibgda_post_send_db(mlx5gda_qp_devctx* qp) { uint32_t num_posted = static_cast(qp->wq_head); // DBR write — always done (NIC polls doorbell record in GPU memory) +#if defined(MOONCAKE_EP_USE_MUSA) && __MUSA_ARCH__ == 310 + // Match mtshmem's IBGDA publish sequence. The bypass atomic exchange makes + // the GPU-memory doorbell record visible to the NIC rather than retaining + // the write in the GPU cache hierarchy. + const uint32_t dbr_value = mc_bswap32(num_posted); + __threadfence_system_noflush(); + asm volatile("LSU.ATOM_32.XCHG %0, %1, slc=byp;" + : + : "R"(reinterpret_cast(&qp->dbr->send_counter)), + "R"(dbr_value)); +#else mc_st_release_u32(reinterpret_cast(&qp->dbr->send_counter), mc_bswap32(num_posted)); +#endif // BF (Blue Flame) doorbell — only if BF register is mapped into GPU VA. - // On MUSA, musaHostRegisterIoMemory fails for MMIO addresses, so bf is - // NULL and we rely on DBR-only mode (slightly higher latency). if (qp->bf != nullptr) { +#if defined(MOONCAKE_EP_USE_MUSA) && __MUSA_ARCH__ == 310 + struct { + __be32 opmod_idx_opcode; + __be32 qpn_ds; + } doorbell{ + mc_bswap32(num_posted << 8), + mc_bswap32(qp->qpn << 8), + }; + __threadfence_system_noflush(); + asm volatile("LSU.ATOM_64.XCHG %0, %1, slc=byp;" + : + : "R"(reinterpret_cast(qp->bf)), + "R"(*reinterpret_cast(&doorbell))); +#else auto* last_wqe = qp->wq + ((num_posted - 1) & qp->wqeid_mask); mc_st_release_u64(reinterpret_cast(qp->bf + qp->bf_offset), *reinterpret_cast(last_wqe)); qp->bf_offset ^= MLX5GDA_BF_SIZE; +#endif } } diff --git a/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp b/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp index 7397ede7d5..8edd6b90a0 100644 --- a/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp +++ b/mooncake-transfer-engine/src/transport/device/ibgda_device_transport.cpp @@ -26,6 +26,7 @@ #include #include +#include #include #include #include @@ -42,11 +43,14 @@ namespace mooncake { namespace device { static constexpr size_t kCtrlBufSize = 1024ULL * 1024 * 1024; // 1 GiB +static constexpr size_t kHostMappedCqMinSize = 256ULL * 1024 * 1024; +static constexpr size_t kIbgdaQpWqebbCount = 16384; enum class ControlMemoryMode { kGpuDmabuf, kGpuVa, kHostMapped, + kPerQpGpuVaHostMappedCq, }; static const char* controlMemoryModeName(ControlMemoryMode mode) { @@ -57,6 +61,8 @@ static const char* controlMemoryModeName(ControlMemoryMode mode) { return "gpu-va"; case ControlMemoryMode::kHostMapped: return "host-mapped"; + case ControlMemoryMode::kPerQpGpuVaHostMappedCq: + return "per-qp-gpu-va-host-mapped-cq"; } return "unknown"; } @@ -256,23 +262,53 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { "trying GPU-VA control buffer"; return allocateGpuVaOrHostControlBuffer(); #endif +#if defined(USE_MUSA) + return allocatePerQpGpuVaControlBuffers(); +#else return allocateControlBuffer(ControlMemoryMode::kGpuVa); +#endif } int createQueuePairs(void* stream_ptr) override { auto stream = static_cast(stream_ptr); - mlx5gda_control_region_allocator region_allocator{ + mlx5gda_control_region_allocator dmabuf_region_allocator{ .context = this, .allocate = allocateDmabufControlRegionThunk, .release = releaseDmabufControlRegionThunk, }; - const mlx5gda_control_region_allocator* allocator = - ctrl_buf_mode_ == ControlMemoryMode::kGpuDmabuf ? ®ion_allocator - : nullptr; + mlx5gda_control_region_allocator gpu_va_region_allocator{ + .context = this, + .allocate = allocateGpuVaControlRegionThunk, + .release = releaseGpuVaControlRegionThunk, + }; + const mlx5gda_control_region_allocator* cq_region_allocator = + ctrl_buf_mode_ == ControlMemoryMode::kGpuDmabuf + ? &dmabuf_region_allocator + : nullptr; + const mlx5gda_control_region_allocator* qp_region_allocator = + ctrl_buf_mode_ == ControlMemoryMode::kGpuDmabuf + ? &dmabuf_region_allocator + : ctrl_buf_mode_ == ControlMemoryMode::kPerQpGpuVaHostMappedCq + ? &gpu_va_region_allocator + : nullptr; + const mlx5gda_control_buffer ctrl{ + .addr = ctrl_buf_, + .dev_addr = ctrl_buf_dev_, + .umem = ctrl_buf_umem_, + .heap = ctrl_buf_heap_, + }; + const mlx5gda_control_buffer cq = + cq_buf_ ? mlx5gda_control_buffer{ + .addr = cq_buf_, + .dev_addr = cq_buf_dev_, + .umem = cq_buf_umem_, + .heap = cq_buf_heap_, + } + : ctrl; for (int i = 0; i < num_qps_; ++i) { mlx5gda_qp* qp = mlx5gda_create_rc_qp( - mpd_, ctrl_buf_, ctrl_buf_umem_, ctrl_buf_heap_, pd_, 16384, 1, - stream, allocator); + mpd_, &ctrl, &cq, pd_, kIbgdaQpWqebbCount, 1, stream, + cq_region_allocator, qp_region_allocator); if (!qp) { const int qp_errno = errno; LOG(ERROR) << "[EP IBGDA] mlx5gda_create_rc_qp failed at " << i; @@ -282,31 +318,18 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { } if (mlx5gda_modify_rc_qp_rst2init(qp, 0)) { LOG(ERROR) << "[EP IBGDA] rst2init failed at " << i; - mlx5gda_destroy_qp(ctrl_buf_heap_, qp); + mlx5gda_destroy_qp(qp); return -1; } cudaStreamSynchronize(stream); - const bool split_regions = - ctrl_buf_mode_ == ControlMemoryMode::kGpuDmabuf; mlx5gda_qp_devctx devctx{ .qpn = qp->qpn, .wqeid_mask = qp->num_wqebb - 1, .mutex = 0, - .wq = split_regions ? reinterpret_cast(qp->wq) - : reinterpret_cast( - static_cast(ctrl_buf_dev_) + - qp->wq_offset), - .cq = split_regions - ? reinterpret_cast(qp->send_cq->cq_buf) - : reinterpret_cast( - static_cast(ctrl_buf_dev_) + - qp->send_cq->cq_offset), - .dbr = split_regions - ? reinterpret_cast(qp->dbr) - : reinterpret_cast( - static_cast(ctrl_buf_dev_) + - qp->dbr_offset), - .bf = static_cast(qp->uar->reg_addr), + .wq = reinterpret_cast(qp->dev_wq), + .cq = reinterpret_cast(qp->send_cq->dev_cq_buf), + .dbr = reinterpret_cast(qp->dev_dbr), + .bf = qp->bf_device_addr, .bf_offset = 0, .wq_head = 0, .wq_tail = 0, @@ -492,7 +515,7 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { LOG(ERROR) << "[EP IBGDA] cudaMalloc failed for DMA-BUF control " "region size=" << size << ": " << cudaGetErrorString(cuda_error); - errno = cuda_error == cudaErrorMemoryAllocation ? ENOMEM : EIO; + errno = EIO; return -1; } @@ -565,6 +588,54 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { *region = {}; } + int allocateGpuVaControlRegion(size_t requested_size, + mlx5gda_control_region* region) { + if (!region || requested_size == 0) { + errno = EINVAL; + return -1; + } + + const long page_size = sysconf(_SC_PAGESIZE); + if (page_size <= 0) { + errno = EIO; + return -1; + } + const size_t page_mask = static_cast(page_size) - 1; + const size_t size = (requested_size + page_mask) & ~page_mask; + + void* addr = nullptr; + const cudaError_t cuda_error = cudaMalloc(&addr, size); + if (cuda_error != cudaSuccess) { + LOG(ERROR) << "[EP IBGDA] cudaMalloc failed for per-QP GPU-VA " + "control region size=" + << size << ": " << cudaGetErrorString(cuda_error); + errno = EIO; + return -1; + } + + mlx5dv_devx_umem* umem = + mlx5dv_devx_umem_reg(ctx_, addr, size, IBV_ACCESS_LOCAL_WRITE); + if (!umem) { + PLOG(ERROR) << "[EP IBGDA] Control UMEM registration failed for " + "per-QP GPU-VA region"; + cudaFree(addr); + return -1; + } + *region = mlx5gda_control_region{ + .addr = addr, + .size = size, + .umem = umem, + }; + return 0; + } + + void releaseGpuVaControlRegion(mlx5gda_control_region* region) { + if (!region) return; + if (region->umem) mlx5dv_devx_umem_dereg(region->umem); + if (region->addr) cudaFree(region->addr); + *region = {}; + } + static int allocateDmabufControlRegionThunk( void* context, size_t size, mlx5gda_control_region* region) { return static_cast(context) @@ -577,6 +648,18 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { ->releaseDmabufControlRegion(region); } + static int allocateGpuVaControlRegionThunk(void* context, size_t size, + mlx5gda_control_region* region) { + return static_cast(context) + ->allocateGpuVaControlRegion(size, region); + } + + static void releaseGpuVaControlRegionThunk(void* context, + mlx5gda_control_region* region) { + static_cast(context) + ->releaseGpuVaControlRegion(region); + } + int allocateControlBuffer(ControlMemoryMode mode) { ctrl_buf_mode_ = mode; if (mode == ControlMemoryMode::kHostMapped) { @@ -644,6 +727,79 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { return 0; } + int allocatePerQpGpuVaControlBuffers() { + ctrl_buf_mode_ = ControlMemoryMode::kPerQpGpuVaHostMappedCq; + if (allocateHostMappedCqBuffer() != 0) { + freeControlBuffer(); + return -1; + } + LOG(INFO) << "[EP IBGDA] Using per-QP GPU WQ/DBR regions and " + "host-backed mapped CQ buffer"; + return 0; + } + + int allocateHostMappedCqBuffer() { + // Each QP owns a 16K-entry CQ. A CQ consumes exactly 1 MiB, but its + // trailing DBR makes the next CQ start on a fresh adapter page. The + // fixed 256 MiB baseline therefore cannot hold 256 QPs. + constexpr size_t kAdapterPageSize = 4096; + constexpr size_t kCqBytesPerQp = + kIbgdaQpWqebbCount * sizeof(mlx5_cqe64) + kAdapterPageSize; + const size_t cq_buf_size = + std::max(kHostMappedCqMinSize, + static_cast(num_qps_) * kCqBytesPerQp); + void* ptr = nullptr; + const int ret = posix_memalign(&ptr, 4096, cq_buf_size); + if (ret != 0) { + LOG(ERROR) << "[EP IBGDA] posix_memalign cq_buf failed: " << ret; + return -1; + } + cq_buf_ = ptr; + std::memset(cq_buf_, 0, cq_buf_size); + cudaError_t err = + cudaHostRegister(cq_buf_, cq_buf_size, + cudaHostRegisterPortable | cudaHostRegisterMapped); + if (err != cudaSuccess) { + LOG(ERROR) << "[EP IBGDA] cudaHostRegister cq_buf failed: " + << cudaGetErrorString(err); + free(cq_buf_); + cq_buf_ = nullptr; + return -1; + } + err = cudaHostGetDevicePointer(&cq_buf_dev_, cq_buf_, 0); + if (err != cudaSuccess) { + LOG(ERROR) << "[EP IBGDA] cudaHostGetDevicePointer cq_buf failed: " + << cudaGetErrorString(err); + cudaHostUnregister(cq_buf_); + free(cq_buf_); + cq_buf_ = nullptr; + return -1; + } + cq_buf_umem_ = mlx5dv_devx_umem_reg(ctx_, cq_buf_, cq_buf_size, + IBV_ACCESS_LOCAL_WRITE); + if (!cq_buf_umem_) { + LOG(ERROR) << "[EP IBGDA] CQ UMEM registration failed (errno=" + << errno << ")"; + cudaHostUnregister(cq_buf_); + free(cq_buf_); + cq_buf_ = nullptr; + cq_buf_dev_ = nullptr; + return -1; + } + cq_buf_heap_ = memheap_create(cq_buf_size); + if (!cq_buf_heap_) { + LOG(ERROR) << "[EP IBGDA] memheap_create cq_buf failed"; + mlx5dv_devx_umem_dereg(cq_buf_umem_); + cq_buf_umem_ = nullptr; + cudaHostUnregister(cq_buf_); + free(cq_buf_); + cq_buf_ = nullptr; + cq_buf_dev_ = nullptr; + return -1; + } + return 0; + } + int allocateGpuVaOrHostControlBuffer() { if (allocateControlBuffer(ControlMemoryMode::kGpuVa) == 0) return 0; @@ -657,6 +813,7 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { const bool using_dmabuf_control_regions = ctrl_buf_mode_ == ControlMemoryMode::kGpuDmabuf; if (ctrl_buf_mode_ == ControlMemoryMode::kHostMapped || + ctrl_buf_mode_ == ControlMemoryMode::kPerQpGpuVaHostMappedCq || (!isCreateQpBadParam(failure) && !using_dmabuf_control_regions)) return false; @@ -685,12 +842,26 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { void destroyQueuePairs() { for (auto* qp : qps_) { - if (qp) mlx5gda_destroy_qp(ctrl_buf_heap_, qp); + if (qp) mlx5gda_destroy_qp(qp); } qps_.clear(); } void freeControlBuffer() { + if (cq_buf_heap_) { + memheap_destroy(cq_buf_heap_); + cq_buf_heap_ = nullptr; + } + if (cq_buf_umem_) { + mlx5dv_devx_umem_dereg(cq_buf_umem_); + cq_buf_umem_ = nullptr; + } + if (cq_buf_) { + cudaHostUnregister(cq_buf_); + free(cq_buf_); + cq_buf_ = nullptr; + cq_buf_dev_ = nullptr; + } if (ctrl_buf_heap_) { memheap_destroy(ctrl_buf_heap_); ctrl_buf_heap_ = nullptr; @@ -760,6 +931,10 @@ class IbgdaDeviceTransportImpl : public RdmaTransport { ControlMemoryMode ctrl_buf_mode_ = ControlMemoryMode::kGpuVa; mlx5dv_devx_umem* ctrl_buf_umem_ = nullptr; memheap* ctrl_buf_heap_ = nullptr; + void* cq_buf_ = nullptr; + void* cq_buf_dev_ = nullptr; + mlx5dv_devx_umem* cq_buf_umem_ = nullptr; + memheap* cq_buf_heap_ = nullptr; // QPs std::vector qps_; diff --git a/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp b/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp index 4e85aaf489..c48396322d 100644 --- a/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp +++ b/mooncake-transfer-engine/src/transport/device/mlx5gda.cpp @@ -1,4 +1,6 @@ #include +#include +#include #include "cuda_alike.h" @@ -31,11 +33,7 @@ constexpr T round_up_pow2(T n) { #define IBGDA_ROUND_UP_POW2_OR_0(_n) (((_n) == 0) ? 0 : round_up_pow2(_n)) static void print_cuda_error(const char* msg) { -#ifdef USE_MUSA - const char* err_str = musaGetErrorString(musaGetLastError()); -#else const char* err_str = cudaGetErrorString(cudaGetLastError()); -#endif fprintf(stderr, "%s: %s\n", msg, err_str); } @@ -81,14 +79,6 @@ static void print_devx_create_qp_failure(const uint8_t* cmd_out, dbr_offset); } -// Create UAR for BF (Blue Flame) doorbell ringing. -// On CUDA: registers the BF MMIO region into GPU address space so the -// GPU kernel can directly write the doorbell (lowest latency). -// On MUSA: musaHostRegisterIoMemory is not supported for MMIO addresses, -// so we skip BF registration and return a UAR with reg_addr=NULL. -// The GPU kernel will use DBR-only mode (write to memory-mapped -// doorbell record, NIC polls it) — slightly higher latency but -// functionally correct. static struct mlx5dv_devx_uar* create_uar(struct ibv_context* ctx) { struct mlx5dv_devx_uar* uar = mlx5dv_devx_alloc_uar(ctx, MLX5DV_UAR_ALLOC_TYPE_BF); @@ -96,25 +86,35 @@ static struct mlx5dv_devx_uar* create_uar(struct ibv_context* ctx) { errno = EIO; return NULL; } -#ifdef USE_MUSA - // MUSA cannot map MMIO addresses into GPU VA. Skip the - // musaHostRegister(IoMemory) call entirely — attempting it - // corrupts the MUSA runtime, causing all subsequent device-side - // fill operations to fail with "illegal memory access". - // Use DBR-only mode: the kernel writes to the doorbell record - // in GPU memory instead of the BF MMIO register. - uar->reg_addr = NULL; -#else - if (cudaHostRegister(uar->reg_addr, MLX5GDA_BF_SIZE * 2, - cudaHostRegisterPortable | cudaHostRegisterMapped | - cudaHostRegisterIoMemory) != cudaSuccess) { - print_cuda_error("Failed to register MMIO memory"); + return uar; +} + +static int map_uar_for_device(struct mlx5dv_devx_uar* uar, + struct mlx5gda_qp* qp) { + const cudaError_t register_result = + cudaHostRegister(uar->reg_addr, MLX5GDA_BF_SIZE * 2, + cudaHostRegisterMapped | cudaHostRegisterIoMemory); + void* device_base = nullptr; + if (cudaHostGetDevicePointer(&device_base, uar->reg_addr, 0) != + cudaSuccess) { + if (register_result == cudaSuccess) cudaHostUnregister(uar->reg_addr); + print_cuda_error("Failed to map BF MMIO page"); errno = EIO; - mlx5dv_devx_free_uar(uar); - return NULL; + return -1; } -#endif - return uar; + qp->bf_host_base = uar->reg_addr; + qp->bf_device_addr = static_cast(device_base); + qp->bf_host_registration_owner = register_result == cudaSuccess; + return 0; +} + +static void unmap_uar_for_device(struct mlx5gda_qp* qp) { + if (qp->bf_host_registration_owner) { + cudaHostUnregister(qp->bf_host_base); + qp->bf_host_registration_owner = false; + } + qp->bf_host_base = nullptr; + qp->bf_device_addr = nullptr; } static bool uses_control_regions( @@ -133,19 +133,15 @@ static void release_control_region( static void destroy_uar(struct mlx5dv_devx_uar* uar) { if (!uar) return; - if (uar->reg_addr) { - if (cudaHostUnregister(uar->reg_addr) != cudaSuccess) { - print_cuda_error("Failed to unregister MMIO memory"); - } - } mlx5dv_devx_free_uar(uar); } struct mlx5gda_cq* mlx5gda_create_cq( void* ctrl_buf, struct mlx5dv_devx_umem* ctrl_buf_umem, struct memheap* ctrl_buf_heap, struct ibv_pd* pd, int cqe, - cudaStream_t stream, - const struct mlx5gda_control_region_allocator* region_allocator) { + cudaStream_t stream, void* cq_buf_arg, void* cq_buf_dev, + struct mlx5dv_devx_umem* cq_buf_umem, struct memheap* cq_buf_heap, + const struct mlx5gda_control_region_allocator* cq_region_allocator) { struct mlx5gda_cq* cq = NULL; struct mlx5dv_devx_uar* uar = NULL; uint32_t eqn = 0; @@ -153,15 +149,17 @@ struct mlx5gda_cq* mlx5gda_create_cq( size_t dbr_offset = -1; struct mlx5dv_devx_obj* mlx5_cq = NULL; uint32_t cqn = 0; - void* cq_buf = ctrl_buf; - void* dbr = ctrl_buf; - struct mlx5dv_devx_umem* cq_umem = ctrl_buf_umem; - struct mlx5dv_devx_umem* dbr_umem = ctrl_buf_umem; - const bool split_regions = uses_control_regions(region_allocator); + void* cq_buf = cq_buf_arg ? cq_buf_arg : ctrl_buf; + void* dbr = cq_buf; + struct mlx5dv_devx_umem* cq_umem = + cq_buf_umem ? cq_buf_umem : ctrl_buf_umem; + struct mlx5dv_devx_umem* dbr_umem = cq_umem; + struct memheap* cq_heap = cq_buf_heap ? cq_buf_heap : ctrl_buf_heap; + const bool split_regions = uses_control_regions(cq_region_allocator); struct ibv_context* ctx = pd->context; void* cq_context = NULL; - bool ctrl_host = !split_regions && is_host_control_buffer(ctrl_buf); + bool ctrl_host = !split_regions && is_host_control_buffer(cq_buf); if (cqe <= 0) { errno = EINVAL; @@ -176,16 +174,16 @@ struct mlx5gda_cq* mlx5gda_create_cq( if (!cq) goto fail; if (split_regions) { - cq->region_allocator = *region_allocator; - if (region_allocator->allocate(region_allocator->context, - num_cqe * sizeof(struct mlx5_cqe64), - &cq->cq_region) != 0) { + cq->region_allocator = *cq_region_allocator; + if (cq_region_allocator->allocate(cq_region_allocator->context, + num_cqe * sizeof(struct mlx5_cqe64), + &cq->cq_region) != 0) { perror("Failed to allocate CQ control region"); goto fail; } - if (region_allocator->allocate(region_allocator->context, - sizeof(struct mlx5gda_cq_dbr), - &cq->dbr_region) != 0) { + if (cq_region_allocator->allocate(cq_region_allocator->context, + sizeof(struct mlx5gda_cq_dbr), + &cq->dbr_region) != 0) { perror("Failed to allocate CQ DBR control region"); goto fail; } @@ -196,15 +194,14 @@ struct mlx5gda_cq* mlx5gda_create_cq( cq_offset = 0; dbr_offset = 0; } else { - cq_offset = memheap_aligned_alloc(ctrl_buf_heap, - num_cqe * sizeof(struct mlx5_cqe64), - (size_t)1 << MLX5_ADAPTER_PAGE_SHIFT); + cq_offset = + memheap_aligned_alloc(cq_heap, num_cqe * sizeof(struct mlx5_cqe64), + (size_t)1 << MLX5_ADAPTER_PAGE_SHIFT); if (cq_offset == (size_t)-1) { perror("Failed to allocate CQ memory"); goto fail; } - dbr_offset = - memheap_alloc(ctrl_buf_heap, sizeof(struct mlx5gda_cq_dbr)); + dbr_offset = memheap_alloc(cq_heap, sizeof(struct mlx5gda_cq_dbr)); if (dbr_offset == (size_t)-1) { perror("Failed to allocate CQ DBR memory"); goto fail; @@ -271,7 +268,12 @@ struct mlx5gda_cq* mlx5gda_create_cq( cq->cq_offset = cq_offset; cq->dbr_offset = dbr_offset; + cq->dev_base = cq_buf_dev ? cq_buf_dev : ctrl_buf; + cq->heap = cq_heap; cq->cq_buf = static_cast(cq_buf) + cq_offset; + cq->dev_cq_buf = split_regions + ? cq->cq_buf + : static_cast(cq->dev_base) + cq_offset; cq->dbr = static_cast(dbr) + dbr_offset; cq->cqe = num_cqe; cq->cqn = cqn; @@ -287,14 +289,14 @@ struct mlx5gda_cq* mlx5gda_create_cq( free(cq); } if (!split_regions) { - if (cq_offset != (size_t)-1) memheap_free(ctrl_buf_heap, cq_offset); - if (dbr_offset != (size_t)-1) memheap_free(ctrl_buf_heap, dbr_offset); + if (cq_offset != (size_t)-1) memheap_free(cq_heap, cq_offset); + if (dbr_offset != (size_t)-1) memheap_free(cq_heap, dbr_offset); } errno = saved_errno; return NULL; } -void mlx5gda_destroy_cq(struct memheap* ctrl_buf_heap, struct mlx5gda_cq* cq) { +void mlx5gda_destroy_cq(struct mlx5gda_cq* cq) { if (!cq) return; if (cq->mcq) { mlx5dv_devx_obj_destroy(cq->mcq); @@ -306,17 +308,18 @@ void mlx5gda_destroy_cq(struct memheap* ctrl_buf_heap, struct mlx5gda_cq* cq) { release_control_region(&cq->region_allocator, &cq->dbr_region); release_control_region(&cq->region_allocator, &cq->cq_region); } else { - memheap_free(ctrl_buf_heap, cq->cq_offset); - memheap_free(ctrl_buf_heap, cq->dbr_offset); + memheap_free(cq->heap, cq->cq_offset); + memheap_free(cq->heap, cq->dbr_offset); } free(cq); } struct mlx5gda_qp* mlx5gda_create_rc_qp( - struct mlx5dv_pd mpd, void* ctrl_buf, - struct mlx5dv_devx_umem* ctrl_buf_umem, struct memheap* ctrl_buf_heap, - struct ibv_pd* pd, int wqe, uint8_t port_num, cudaStream_t stream, - const struct mlx5gda_control_region_allocator* region_allocator) { + struct mlx5dv_pd mpd, const struct mlx5gda_control_buffer* ctrl, + const struct mlx5gda_control_buffer* cq, struct ibv_pd* pd, int wqe, + uint8_t port_num, cudaStream_t stream, + const struct mlx5gda_control_region_allocator* cq_region_allocator, + const struct mlx5gda_control_region_allocator* qp_region_allocator) { mlx5gda_reset_create_qp_failure(); struct mlx5gda_qp* qp = NULL; @@ -325,17 +328,17 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( struct mlx5dv_devx_obj* mlx5_qp = NULL; size_t wq_offset = -1; size_t dbr_offset = -1; - void* wq = ctrl_buf; - void* dbr = ctrl_buf; - struct mlx5dv_devx_umem* wq_umem = ctrl_buf_umem; - struct mlx5dv_devx_umem* dbr_umem = ctrl_buf_umem; - const bool split_regions = uses_control_regions(region_allocator); + void* wq = ctrl->addr; + void* dbr = ctrl->addr; + struct mlx5dv_devx_umem* wq_umem = ctrl->umem; + struct mlx5dv_devx_umem* dbr_umem = ctrl->umem; + const bool split_regions = uses_control_regions(qp_region_allocator); struct ibv_context* ctx = pd->context; void* qp_context = NULL; void* cap = NULL; uint32_t cqe_version = 0; - bool ctrl_host = !split_regions && is_host_control_buffer(ctrl_buf); + bool ctrl_host = !split_regions && is_host_control_buffer(ctrl->addr); if (wqe <= 0) { errno = EINVAL; @@ -354,7 +357,7 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( perror("Failed to allocate QP memory"); goto fail; } - if (split_regions) qp->region_allocator = *region_allocator; + if (split_regions) qp->region_allocator = *qp_region_allocator; qp->port_num = port_num; if (ibv_query_port(ctx, port_num, &qp->port_attr) != 0) { @@ -382,8 +385,9 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( } // Create send_cq on GPU memory. - send_cq = mlx5gda_create_cq(ctrl_buf, ctrl_buf_umem, ctrl_buf_heap, pd, wqe, - stream, region_allocator); + send_cq = mlx5gda_create_cq(ctrl->addr, ctrl->umem, ctrl->heap, pd, wqe, + stream, cq->addr, cq->dev_addr, cq->umem, + cq->heap, cq_region_allocator); if (send_cq == NULL) { perror("mlx5gda_create_cq failed"); goto fail; @@ -394,17 +398,19 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( perror("Failed to create UAR"); goto fail; } + if (map_uar_for_device(uar, qp) != 0) goto fail; if (split_regions) { - if (region_allocator->allocate(region_allocator->context, - num_wqebb * sizeof(struct mlx5gda_wqebb), - &qp->wq_region) != 0) { + if (qp_region_allocator->allocate( + qp_region_allocator->context, + num_wqebb * sizeof(struct mlx5gda_wqebb), + &qp->wq_region) != 0) { perror("Failed to allocate WQ control region"); goto fail; } - if (region_allocator->allocate(region_allocator->context, - sizeof(struct mlx5gda_wq_dbr), - &qp->dbr_region) != 0) { + if (qp_region_allocator->allocate(qp_region_allocator->context, + sizeof(struct mlx5gda_wq_dbr), + &qp->dbr_region) != 0) { perror("Failed to allocate QP DBR control region"); goto fail; } @@ -416,15 +422,14 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( dbr_offset = 0; } else { wq_offset = memheap_aligned_alloc( - ctrl_buf_heap, num_wqebb * sizeof(struct mlx5gda_wqebb), + ctrl->heap, num_wqebb * sizeof(struct mlx5gda_wqebb), (size_t)1 << MLX5_ADAPTER_PAGE_SHIFT); if (wq_offset == (size_t)-1) { perror("Failed to allocate WQ memory"); goto fail; } - dbr_offset = - memheap_alloc(ctrl_buf_heap, sizeof(struct mlx5gda_wq_dbr)); + dbr_offset = memheap_alloc(ctrl->heap, sizeof(struct mlx5gda_wq_dbr)); if (dbr_offset == (size_t)-1) { perror("Failed to allocate DBR memory"); goto fail; @@ -496,8 +501,14 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( qp->num_wqebb = num_wqebb; qp->wq_offset = wq_offset; qp->dbr_offset = dbr_offset; + qp->wq_heap = ctrl->heap; qp->wq = static_cast(wq) + wq_offset; qp->dbr = static_cast(dbr) + dbr_offset; + qp->dev_wq = + split_regions ? qp->wq : static_cast(ctrl->dev_addr) + wq_offset; + qp->dev_dbr = split_regions + ? qp->dbr + : static_cast(ctrl->dev_addr) + dbr_offset; return qp; fail: @@ -506,10 +517,11 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( mlx5dv_devx_obj_destroy(mlx5_qp); } if (uar) { + if (qp) unmap_uar_for_device(qp); destroy_uar(uar); } if (send_cq) { - mlx5gda_destroy_cq(ctrl_buf_heap, send_cq); + mlx5gda_destroy_cq(send_cq); } if (qp) { release_control_region(&qp->region_allocator, &qp->dbr_region); @@ -519,34 +531,35 @@ struct mlx5gda_qp* mlx5gda_create_rc_qp( free(qp); } if (!split_regions && wq_offset != (size_t)-1) { - memheap_free(ctrl_buf_heap, wq_offset); + memheap_free(ctrl->heap, wq_offset); } if (!split_regions && dbr_offset != (size_t)-1) { - memheap_free(ctrl_buf_heap, dbr_offset); + memheap_free(ctrl->heap, dbr_offset); } errno = saved_errno; return NULL; } -void mlx5gda_destroy_qp(struct memheap* ctrl_buf_heap, struct mlx5gda_qp* qp) { +void mlx5gda_destroy_qp(struct mlx5gda_qp* qp) { if (qp->mqp) { mlx5dv_devx_obj_destroy(qp->mqp); } if (qp->uar) { + unmap_uar_for_device(qp); destroy_uar(qp->uar); } if (qp->send_cq) { - mlx5gda_destroy_cq(ctrl_buf_heap, qp->send_cq); + mlx5gda_destroy_cq(qp->send_cq); } if (uses_control_regions(&qp->region_allocator)) { release_control_region(&qp->region_allocator, &qp->dbr_region); release_control_region(&qp->region_allocator, &qp->wq_region); } else { if (qp->wq_offset != (size_t)-1) { - memheap_free(ctrl_buf_heap, qp->wq_offset); + memheap_free(qp->wq_heap, qp->wq_offset); } if (qp->dbr_offset != (size_t)-1) { - memheap_free(ctrl_buf_heap, qp->dbr_offset); + memheap_free(qp->wq_heap, qp->dbr_offset); } } if (qp) { From 69705bab8ec221258db4d0264a9bcbcb38ee1ffa Mon Sep 17 00:00:00 2001 From: Hubert Zhu <55722581+Hubert-Zhu@users.noreply.github.com> Date: Sun, 9 Aug 2026 23:00:01 -0400 Subject: [PATCH 011/483] [Store] Remove dead DRAM/NoF metric overload declarations (#3344) These parameterless inc/dec overloads were declared in the header but never implemented or called; keep the segment-keyed APIs that actually exist. Co-authored-by: Cursor --- mooncake-store/include/master_metric_manager.h | 8 -------- 1 file changed, 8 deletions(-) diff --git a/mooncake-store/include/master_metric_manager.h b/mooncake-store/include/master_metric_manager.h index c98c43b177..8e4cd4903a 100644 --- a/mooncake-store/include/master_metric_manager.h +++ b/mooncake-store/include/master_metric_manager.h @@ -81,10 +81,6 @@ class MasterMetricManager { CacheHitStatDict calculate_cache_stats(); // Memory Storage Metrics - void inc_allocated_mem_size(int64_t val = 1); - void dec_allocated_mem_size(int64_t val = 1); - void inc_total_mem_capacity(int64_t val = 1); - void dec_total_mem_capacity(int64_t val = 1); int64_t get_allocated_mem_size(); int64_t get_total_mem_capacity(); double get_segment_mem_used_ratio(const std::string& segment); @@ -99,10 +95,6 @@ class MasterMetricManager { void remove_segment_metrics(const std::string& segment); // NoF segment Metrics - void inc_allocated_nof_size(int64_t val = 1); - void dec_allocated_nof_size(int64_t val = 1); - void inc_total_nof_capacity(int64_t val = 1); - void dec_total_nof_capacity(int64_t val = 1); int64_t get_allocated_nof_size(); int64_t get_total_nof_capacity(); double get_segment_nof_used_ratio(const std::string& segment); From d4ba937e1d9d5c8bdae1cd5626a5eef54d823a10 Mon Sep 17 00:00:00 2001 From: Aoi Date: Mon, 10 Aug 2026 12:19:01 +0800 Subject: [PATCH 012/483] [Store] Remove transient MasterService config fields (#3313) --- mooncake-store/include/master_service.h | 25 +----- .../include/master_snapshot_manager.h | 7 -- mooncake-store/src/master_service.cpp | 89 +++++++------------ .../src/master_snapshot_manager.cpp | 8 +- .../master_service_test_for_snapshot_base.h | 11 +-- .../snapshot/snapshot_child_process_test.cpp | 57 ++++++------ 6 files changed, 62 insertions(+), 135 deletions(-) diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index 2988f21ba6..93cea6fe6a 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -875,7 +875,8 @@ class MasterService { void setHttpMetadataRemoteUrl(const std::string& metadata_connstring); private: - std::unique_ptr CreateSnapshotCatalogStore(); + std::unique_ptr CreateSnapshotCatalogStore( + const MasterServiceConfig& config); // Restore master state void RestoreState(); @@ -2085,9 +2086,6 @@ class MasterService { // from any GetReplicaList caller without additional locking. std::unique_ptr promotion_sketch_; - const std::string ha_backend_type_; - - const std::string ha_backend_connstring_; const bool enable_oplog_; const uint32_t oplog_batch_max_entries_; @@ -2095,14 +2093,10 @@ class MasterService { const std::string cluster_id_; // root filesystem directory for persistent storage const std::string root_fs_dir_; - // global 3fs/nfs segment size - int64_t global_file_segment_size_; // storage backend eviction configuration const bool enable_disk_eviction_; const uint64_t quota_bytes_; const bool enable_multi_tenants_; - const std::string tenant_quota_connector_type_; - const std::string tenant_quota_connector_uri_; std::unique_ptr tenant_quota_policy_store_; mutable std::mutex tenant_quota_policy_mutex_; mutable std::mutex tenant_quota_recompute_mutex_; @@ -2143,17 +2137,6 @@ class MasterService { const AllocationStrategyType allocation_strategy_type_; std::shared_ptr allocation_strategy_; - bool enable_snapshot_restore_ = false; - - bool enable_snapshot_ = false; - std::string snapshot_backup_dir_; - bool use_snapshot_backup_dir_{false}; - uint64_t snapshot_interval_seconds_ = DEFAULT_SNAPSHOT_INTERVAL_SEC; - uint64_t snapshot_child_timeout_seconds_ = - DEFAULT_SNAPSHOT_CHILD_TIMEOUT_SEC; - uint32_t snapshot_retention_count_ = DEFAULT_SNAPSHOT_RETENTION_COUNT; - std::string snapshot_catalog_store_type_{}; - std::string snapshot_catalog_store_connstring_; std::unique_ptr snapshot_object_store_; std::unique_ptr snapshot_catalog_store_; std::unique_ptr snapshot_repository_; @@ -2163,10 +2146,6 @@ class MasterService { // Discarded replicas management const std::chrono::seconds put_start_discard_timeout_sec_; const std::chrono::seconds put_start_release_timeout_sec_; - const std::string cxl_path_; - const size_t cxl_size_; - bool enable_cxl_; - class DiscardedReplicas { public: DiscardedReplicas() = delete; diff --git a/mooncake-store/include/master_snapshot_manager.h b/mooncake-store/include/master_snapshot_manager.h index bd5c04249e..60df9cea99 100644 --- a/mooncake-store/include/master_snapshot_manager.h +++ b/mooncake-store/include/master_snapshot_manager.h @@ -31,18 +31,11 @@ class SnapshotChildProcessTest; } // namespace test struct MasterSnapshotManagerOptions { - bool enable_snapshot{false}; uint64_t snapshot_interval_seconds{0}; uint64_t snapshot_child_timeout_seconds{0}; uint32_t snapshot_retention_count{0}; std::string snapshot_backup_dir; bool use_snapshot_backup_dir{false}; - std::string snapshot_catalog_store_type; - std::string snapshot_catalog_store_connstring; - std::string ha_backend_type; - std::string ha_backend_connstring; - std::string cluster_id; - bool enable_ha{false}; }; /** diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index a201e8cd05..e0e30356c1 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -181,19 +181,14 @@ MasterService::MasterService(const MasterServiceConfig& config) config.nof_heartbeat_failures_threshold), enable_ha_(config.enable_ha), enable_offload_(config.enable_offload), - ha_backend_type_(config.ha_backend_type), - ha_backend_connstring_(config.ha_backend_connstring), enable_oplog_(config.enable_ha && config.enable_oplog && config.ha_backend_type == "etcd"), oplog_batch_max_entries_(config.oplog_batch_max_entries), cluster_id_(config.cluster_id), root_fs_dir_(config.root_fs_dir), - global_file_segment_size_(config.global_file_segment_size), enable_disk_eviction_(config.enable_disk_eviction), quota_bytes_(config.quota_bytes), enable_multi_tenants_(config.enable_multi_tenants), - tenant_quota_connector_type_(config.tenant_quota_connector_type), - tenant_quota_connector_uri_(config.tenant_quota_connector_uri), segment_manager_(config.memory_allocator, config.enable_cxl), nof_segment_manager_(config.memory_allocator), memory_allocator_type_(config.memory_allocator), @@ -201,20 +196,8 @@ MasterService::MasterService(const MasterServiceConfig& config) ? AllocationStrategyType::CXL : config.allocation_strategy_type), allocation_strategy_(CreateAllocationStrategy(allocation_strategy_type_)), - enable_snapshot_restore_(config.enable_snapshot_restore), - enable_snapshot_(config.enable_snapshot), - snapshot_backup_dir_(config.snapshot_backup_dir), - snapshot_interval_seconds_(config.snapshot_interval_seconds), - snapshot_child_timeout_seconds_(config.snapshot_child_timeout_seconds), - snapshot_retention_count_(config.snapshot_retention_count), - snapshot_catalog_store_type_(config.snapshot_catalog_store_type), - snapshot_catalog_store_connstring_( - config.snapshot_catalog_store_connstring), put_start_discard_timeout_sec_(config.put_start_discard_timeout_sec), put_start_release_timeout_sec_(config.put_start_release_timeout_sec), - cxl_path_(config.cxl_path), - cxl_size_(config.cxl_size), - enable_cxl_(config.enable_cxl), offloading_queue_limit_(config.offloading_queue_limit), offload_cap_ratio_(config.offload_cap_ratio), task_manager_(config.task_manager_config) { @@ -232,47 +215,44 @@ MasterService::MasterService(const MasterServiceConfig& config) LOG(INFO) << "Local-first allocation strategy enabled"; } - if (enable_snapshot_ || enable_snapshot_restore_) { + const bool use_snapshot_backup_dir = !config.snapshot_backup_dir.empty(); + if (config.enable_snapshot || config.enable_snapshot_restore) { try { auto object_store_type = ParseSnapshotObjectStoreType(config.snapshot_object_store_type); snapshot_object_store_ = SnapshotObjectStore::Create(object_store_type); - snapshot_catalog_store_ = CreateSnapshotCatalogStore(); + snapshot_catalog_store_ = CreateSnapshotCatalogStore(config); } catch (const std::exception& e) { LOG(ERROR) << "Failed to create snapshot stores: " << e.what(); throw std::runtime_error( fmt::format("Failed to create snapshot stores: {}", e.what())); } - if (!snapshot_backup_dir_.empty()) { - use_snapshot_backup_dir_ = true; - } - // Initialize repository and codec for both save and restore snapshot_repository_ = std::make_unique( snapshot_object_store_.get(), snapshot_catalog_store_.get(), - snapshot_backup_dir_, use_snapshot_backup_dir_); + config.snapshot_backup_dir, use_snapshot_backup_dir); snapshot_codec_ = std::make_unique(); } if (enable_multi_tenants_) { - auto store = CreateTenantQuotaPolicyStore(tenant_quota_connector_type_, - tenant_quota_connector_uri_, - cluster_id_); + auto store = CreateTenantQuotaPolicyStore( + config.tenant_quota_connector_type, + config.tenant_quota_connector_uri, cluster_id_); if (!store) { throw std::invalid_argument(store.error()); } tenant_quota_policy_store_ = std::move(store.value()); } - if (enable_snapshot_restore_) { + if (config.enable_snapshot_restore) { RestoreState(); } if (enable_multi_tenants_) { LoadTenantQuotaPoliciesFromStoreOrThrow(); RebuildTenantQuotaUsageFromMetadata(); } - if (enable_snapshot_ && snapshot_retention_count_ == 0) { + if (config.enable_snapshot && config.snapshot_retention_count == 0) { LOG(ERROR) << "snapshot_retention_count must be greater than 0"; throw std::invalid_argument("snapshot_retention_count must be > 0"); } @@ -398,12 +378,12 @@ MasterService::MasterService(const MasterServiceConfig& config) if (enable_oplog_ && !cluster_id_.empty()) { #ifdef STORE_USE_ETCD - if (ha_backend_connstring_.empty()) { + if (config.ha_backend_connstring.empty()) { LOG(INFO) << "Skipping automatic batch-record OpLog writer " "initialization; no HA backend connstring configured"; } else { ErrorCode connect_err = EtcdHelper::ConnectToEtcdStoreClient( - ha_backend_connstring_.c_str()); + config.ha_backend_connstring.c_str()); if (connect_err != ErrorCode::OK) { throw std::runtime_error(fmt::format( "failed to connect HA batch-record OpLog writer to etcd: " @@ -419,7 +399,7 @@ MasterService::MasterService(const MasterServiceConfig& config) } } #else - if (ha_backend_connstring_.empty()) { + if (config.ha_backend_connstring.empty()) { LOG(INFO) << "Skipping automatic batch-record OpLog writer " "initialization; no HA backend connstring configured"; } else { @@ -465,56 +445,50 @@ MasterService::MasterService(const MasterServiceConfig& config) if (!root_fs_dir_.empty()) { use_disk_replica_ = true; - if (global_file_segment_size_ == std::numeric_limits::max()) { + if (config.global_file_segment_size == + std::numeric_limits::max()) { MasterMetricManager::instance().set_dfs_capacity_unlimited(true); } else { MasterMetricManager::instance().inc_total_file_capacity( - global_file_segment_size_); + config.global_file_segment_size); } } - if (enable_snapshot_ && !enable_oplog_) { + if (config.enable_snapshot && !enable_oplog_) { if (memory_allocator_type_ == BufferAllocatorType::OFFSET) { // Initialize and start snapshot manager MasterSnapshotManagerOptions snapshot_options; - snapshot_options.enable_snapshot = enable_snapshot_; snapshot_options.snapshot_interval_seconds = - snapshot_interval_seconds_; + config.snapshot_interval_seconds; snapshot_options.snapshot_child_timeout_seconds = - snapshot_child_timeout_seconds_; + config.snapshot_child_timeout_seconds; snapshot_options.snapshot_retention_count = - snapshot_retention_count_; - snapshot_options.snapshot_backup_dir = snapshot_backup_dir_; - snapshot_options.use_snapshot_backup_dir = use_snapshot_backup_dir_; - snapshot_options.snapshot_catalog_store_type = - snapshot_catalog_store_type_; - snapshot_options.snapshot_catalog_store_connstring = - snapshot_catalog_store_connstring_; - snapshot_options.ha_backend_type = ha_backend_type_; - snapshot_options.ha_backend_connstring = ha_backend_connstring_; - snapshot_options.cluster_id = cluster_id_; - snapshot_options.enable_ha = enable_ha_; + config.snapshot_retention_count; + snapshot_options.snapshot_backup_dir = config.snapshot_backup_dir; + snapshot_options.use_snapshot_backup_dir = use_snapshot_backup_dir; snapshot_manager_ = std::make_unique( this, snapshot_options, snapshot_mutex_, snapshot_object_store_.get(), snapshot_catalog_store_.get()); snapshot_manager_->Start(); } - } else if (enable_snapshot_ && enable_oplog_) { + } else if (config.enable_snapshot && enable_oplog_) { LOG(INFO) << "Skipping primary snapshot generation in batch-record " "OpLog mode; snapshots are owned by standby"; } - if (enable_cxl_) { + if (config.enable_cxl) { allocation_strategy_ = std::make_shared(); - segment_manager_.initializeCxlAllocator(cxl_path_, cxl_size_); + segment_manager_.initializeCxlAllocator(config.cxl_path, + config.cxl_size); VLOG(1) << "action=start_cxl_global_allocator"; } } std::unique_ptr -MasterService::CreateSnapshotCatalogStore() { - auto catalog_kind = ParseSnapshotCatalogKind(snapshot_catalog_store_type_); +MasterService::CreateSnapshotCatalogStore(const MasterServiceConfig& config) { + auto catalog_kind = + ParseSnapshotCatalogKind(config.snapshot_catalog_store_type); if (!catalog_kind) { throw std::invalid_argument(catalog_kind.error()); } @@ -530,9 +504,10 @@ MasterService::CreateSnapshotCatalogStore() { "redis snapshot catalog store is unavailable in the current " "build"); #else - const auto connstring = !snapshot_catalog_store_connstring_.empty() - ? snapshot_catalog_store_connstring_ - : ha_backend_connstring_; + const auto connstring = + !config.snapshot_catalog_store_connstring.empty() + ? config.snapshot_catalog_store_connstring + : config.ha_backend_connstring; if (connstring.empty()) { throw std::invalid_argument( "redis snapshot catalog store requires a connection " diff --git a/mooncake-store/src/master_snapshot_manager.cpp b/mooncake-store/src/master_snapshot_manager.cpp index 91f2e38609..60582c64cd 100644 --- a/mooncake-store/src/master_snapshot_manager.cpp +++ b/mooncake-store/src/master_snapshot_manager.cpp @@ -108,12 +108,6 @@ void MasterSnapshotManager::SnapshotThreadFunc() { break; } - if (!options_.enable_snapshot) { - // Snapshot is disabled - LOG(INFO) - << "[Snapshot] Snapshot is disabled, waiting for next cycle"; - continue; - } // Fork a child process to save current state std::string snapshot_id = @@ -362,7 +356,7 @@ void MasterSnapshotManager::HandleChildExit(pid_t pid, int status, tl::expected MasterSnapshotManager::ResolveSnapshotSequenceId() const { - if (!options_.enable_ha || !master_service_->enable_oplog_) { + if (!master_service_->enable_ha_ || !master_service_->enable_oplog_) { // OpLog sequence ids start at 1. Returning 0 here is a sentinel that // means "no persisted OpLog boundary", so a standby that later calls // Recover(0) will replay from the first entry when oplog following is diff --git a/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot_base.h b/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot_base.h index 5486e9daf0..013d721e23 100644 --- a/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot_base.h +++ b/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot_base.h @@ -176,20 +176,11 @@ class MasterServiceSnapshotTestBase : public ::testing::Test { EnsureSnapshotStores(service); MasterSnapshotManagerOptions options; - options.enable_snapshot = true; options.snapshot_interval_seconds = 300; options.snapshot_child_timeout_seconds = 300; options.snapshot_retention_count = 3; options.snapshot_backup_dir = ""; options.use_snapshot_backup_dir = false; - options.snapshot_catalog_store_type = - service->snapshot_catalog_store_type_; - options.snapshot_catalog_store_connstring = - service->snapshot_catalog_store_connstring_; - options.ha_backend_type = service->ha_backend_type_; - options.ha_backend_connstring = service->ha_backend_connstring_; - options.cluster_id = service->cluster_id_; - options.enable_ha = service->enable_ha_; auto temp_manager = std::make_unique( service, options, service->snapshot_mutex_, @@ -207,7 +198,7 @@ class MasterServiceSnapshotTestBase : public ::testing::Test { if (!service->snapshot_catalog_store_ && service->snapshot_object_store_) { service->snapshot_catalog_store_ = - service->CreateSnapshotCatalogStore(); + service->CreateSnapshotCatalogStore(MasterServiceConfig{}); } } diff --git a/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp b/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp index 0d9b58769b..62e818000f 100644 --- a/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp +++ b/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp @@ -45,10 +45,16 @@ class SnapshotChildProcessTest : public ::testing::Test { } std::unique_ptr service_; + MasterServiceConfig service_config_; static constexpr const char* kEnvSnapshotLocalPath = "MOONCAKE_SNAPSHOT_LOCAL_PATH"; + void CreateService(MasterServiceConfig config) { + service_config_ = std::move(config); + service_ = std::make_unique(service_config_); + } + void SetUp() override { google::InitGoogleLogging("SnapshotChildProcessTest"); FLAGS_logtostderr = true; @@ -90,7 +96,7 @@ class SnapshotChildProcessTest : public ::testing::Test { .set_snapshot_object_store_type("local") .set_view_version(view_version) .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); } #ifdef STORE_USE_ETCD @@ -112,7 +118,7 @@ class SnapshotChildProcessTest : public ::testing::Test { .set_snapshot_object_store_type("local") .set_view_version(view_version) .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); } void CreateBatchEtcdHASnapshotService(const std::string& cluster_id, @@ -133,7 +139,7 @@ class SnapshotChildProcessTest : public ::testing::Test { .set_snapshot_object_store_type("local") .set_view_version(view_version) .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); } #endif @@ -222,10 +228,6 @@ class SnapshotChildProcessTest : public ::testing::Test { } } - bool GetUseSnapshotBackupDir() { - return service_->use_snapshot_backup_dir_; - } - // Check if a key exists in raw metadata (regardless of replica status) bool KeyExistsInMetadata(MasterService* svc, const std::string& key) { size_t shard_idx = svc->getShardIndex(key); @@ -277,22 +279,15 @@ class SnapshotChildProcessTest : public ::testing::Test { EnsureSnapshotStores(); MasterSnapshotManagerOptions options; - options.enable_snapshot = true; options.snapshot_interval_seconds = - service_->snapshot_interval_seconds_; + service_config_.snapshot_interval_seconds; options.snapshot_child_timeout_seconds = - service_->snapshot_child_timeout_seconds_; - options.snapshot_retention_count = service_->snapshot_retention_count_; - options.snapshot_backup_dir = service_->snapshot_backup_dir_; - options.use_snapshot_backup_dir = service_->use_snapshot_backup_dir_; - options.snapshot_catalog_store_type = - service_->snapshot_catalog_store_type_; - options.snapshot_catalog_store_connstring = - service_->snapshot_catalog_store_connstring_; - options.ha_backend_type = service_->ha_backend_type_; - options.ha_backend_connstring = service_->ha_backend_connstring_; - options.cluster_id = service_->cluster_id_; - options.enable_ha = service_->enable_ha_; + service_config_.snapshot_child_timeout_seconds; + options.snapshot_retention_count = + service_config_.snapshot_retention_count; + options.snapshot_backup_dir = service_config_.snapshot_backup_dir; + options.use_snapshot_backup_dir = + !service_config_.snapshot_backup_dir.empty(); return std::make_unique( service_.get(), options, service_->snapshot_mutex_, @@ -308,7 +303,7 @@ class SnapshotChildProcessTest : public ::testing::Test { if (!service_->snapshot_catalog_store_ && service_->snapshot_object_store_) { service_->snapshot_catalog_store_ = - service_->CreateSnapshotCatalogStore(); + service_->CreateSnapshotCatalogStore(service_config_); } } @@ -598,7 +593,7 @@ TEST_F(SnapshotChildProcessTest, RestoreRebuildsGroupedObjectRouting) { .set_default_kv_lease_ttl(600000) .build(); }; - service_ = std::make_unique(make_config()); + CreateService(make_config()); Segment segment; segment.id = generate_uuid(); @@ -629,7 +624,7 @@ TEST_F(SnapshotChildProcessTest, RestoreRebuildsGroupedObjectRouting) { << "PersistState failed: " << persist_result.error().message; service_.reset(); - service_ = std::make_unique(make_config()); + CreateService(make_config()); auto restored_replicas = service_->GetReplicaList(key, TenantId::Default()); ASSERT_TRUE(restored_replicas.has_value()) @@ -652,7 +647,7 @@ TEST_F(SnapshotChildProcessTest, RestorePreservesObjectChecksum) { .set_default_kv_lease_ttl(600000) .build(); }; - service_ = std::make_unique(make_config()); + CreateService(make_config()); Segment segment; segment.id = generate_uuid(); @@ -681,7 +676,7 @@ TEST_F(SnapshotChildProcessTest, RestorePreservesObjectChecksum) { << "PersistState failed: " << persist_result.error().message; service_.reset(); - service_ = std::make_unique(make_config()); + CreateService(make_config()); auto restored = service_->GetReplicaList(key, TenantId::Default()); ASSERT_TRUE(restored.has_value()); @@ -922,7 +917,7 @@ TEST_F(SnapshotChildProcessTest, RestoreWithoutBackupDir_NoBackupFiles) { auto restore_service = std::make_unique(config); // Step 3: Verify NO backup directory was created - // With empty backup_dir, use_snapshot_backup_dir_ should be false + // With an empty backup_dir, no restore directory should be created. // and no restore directory should exist anywhere in tmp_dir() bool any_restore_dir_found = false; for (auto& entry : fs::recursive_directory_iterator(tmp_dir())) { @@ -1070,7 +1065,7 @@ TEST_F(SnapshotChildProcessTest, RestoreCleansNonCompleteReplica) { .set_snapshot_retention_count(3) .set_snapshot_object_store_type("local") .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); // Mount a segment Segment segment; @@ -1145,7 +1140,7 @@ TEST_F(SnapshotChildProcessTest, RestoreCleansExpiredLease) { .set_snapshot_object_store_type("local") .set_default_kv_lease_ttl(600000) // 10 min lease .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); // Mount a segment Segment segment; @@ -1228,7 +1223,7 @@ TEST_F(SnapshotChildProcessTest, PersistState_FailFast_StopsOnFirstError) { .set_snapshot_retention_count(3) .set_snapshot_object_store_type("local") .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); // Mount a segment to have some data to serialize Segment segment; @@ -1296,7 +1291,7 @@ TEST_F(SnapshotChildProcessTest, UploadFail_WithBackupDir_SavesAllFiles) { .set_snapshot_retention_count(3) .set_snapshot_object_store_type("local") .build(); - service_ = std::make_unique(config); + CreateService(std::move(config)); // Mount a segment to have some data to serialize Segment segment; From fbcb1692a7bd0f9b683819d61738706df25446af Mon Sep 17 00:00:00 2001 From: Schatten Date: Mon, 10 Aug 2026 12:20:06 +0800 Subject: [PATCH 013/483] [Store] Extend MasterScenario DSL for upsert lifecycle (#3328) Signed-off-by: Schatten --- mooncake-store/tests/master_scenario.cpp | 104 ++++++- mooncake-store/tests/master_scenario.h | 177 ++++++++++- mooncake-store/tests/master_scenario_test.cpp | 112 +++++++ .../tests/master_service_scenario_test.cpp | 109 +++++++ mooncake-store/tests/master_service_test.cpp | 278 ------------------ 5 files changed, 479 insertions(+), 301 deletions(-) diff --git a/mooncake-store/tests/master_scenario.cpp b/mooncake-store/tests/master_scenario.cpp index ccf4187bf3..0962ba1ca7 100644 --- a/mooncake-store/tests/master_scenario.cpp +++ b/mooncake-store/tests/master_scenario.cpp @@ -40,10 +40,20 @@ PutStartAction<> PutStart(std::string key, uint64_t size) { return PutStartAction<>(std::move(key), size); } +UpsertStartAction<> UpsertStart(std::string key, uint64_t size) { + return UpsertStartAction<>(std::move(key), size); +} + PutEndAction PutEnd(std::string key) { return {.key = std::move(key)}; } +UpsertEndAction UpsertEnd(std::string key) { return {.key = std::move(key)}; } + PutRevokeAction PutRevoke(std::string key) { return {.key = std::move(key)}; } +UpsertRevokeAction UpsertRevoke(std::string key) { + return {.key = std::move(key)}; +} + RemoveAction Remove(std::string key) { return {.key = std::move(key)}; } ObjectSpec<> Object(std::string key) { return ObjectSpec<>(std::move(key)); } @@ -86,24 +96,25 @@ MasterScenario& MasterScenario::WhenPutStart(PutStartActionData action) { const auto result = service_->PutStart(ActorId(action.actor), action.key, TenantId::Default(), action.size, config); - ValidateActionResult("PutStart(" + action.key + ")", action.expected_error, - result.has_value(), - result ? ErrorCode::OK : result.error()); - if (!result) { + ValidateStartResult("PutStart(" + action.key + ")", action.expected_error, + action.expected_replica_count, + action.expected_replica_status, result); + return *this; +} + +MasterScenario& MasterScenario::WhenUpsertStart(UpsertStartActionData action) { + if (!EnsureService()) { return *this; } - if (action.expected_replica_count.has_value() && - result->size() != *action.expected_replica_count) { - Fail("PutStart(" + action.key + ") returned " + - std::to_string(result->size()) + " replicas; expected " + - std::to_string(*action.expected_replica_count)); - } - if (action.expected_replica_status.has_value() && - std::any_of(result->begin(), result->end(), [&](const auto& replica) { - return replica.status != *action.expected_replica_status; - })) { - Fail("PutStart(" + action.key + ") replica status mismatch"); - } + + ReplicateConfig config; + config.replica_num = 1; + const auto result = + service_->UpsertStart(ActorId(action.actor), action.key, + TenantId::Default(), action.size, config); + ValidateStartResult("UpsertStart(" + action.key + ")", + action.expected_error, action.expected_replica_count, + action.expected_replica_status, result); return *this; } @@ -121,6 +132,20 @@ MasterScenario& MasterScenario::When(PutEndAction action) { return *this; } +MasterScenario& MasterScenario::When(UpsertEndAction action) { + if (!EnsureService()) { + return *this; + } + + const auto result = + service_->UpsertEnd(ActorId(action.actor), action.key, + TenantId::Default(), ReplicaType::MEMORY); + ValidateActionResult("UpsertEnd(" + action.key + ")", action.expected_error, + result.has_value(), + result ? ErrorCode::OK : result.error()); + return *this; +} + MasterScenario& MasterScenario::When(PutRevokeAction action) { if (!EnsureService()) { return *this; @@ -135,6 +160,20 @@ MasterScenario& MasterScenario::When(PutRevokeAction action) { return *this; } +MasterScenario& MasterScenario::When(UpsertRevokeAction action) { + if (!EnsureService()) { + return *this; + } + + const auto result = + service_->UpsertRevoke(ActorId(action.actor), action.key, + TenantId::Default(), ReplicaType::MEMORY); + ValidateActionResult("UpsertRevoke(" + action.key + ")", + action.expected_error, result.has_value(), + result ? ErrorCode::OK : result.error()); + return *this; +} + MasterScenario& MasterScenario::When(RemoveAction action) { if (!EnsureService()) { return *this; @@ -155,6 +194,15 @@ MasterScenario& MasterScenario::ThenObject(ObjectSpecData object, const auto result = service_->GetReplicaList(object.key, TenantId::Default()); + if (expectation == ObjectExpectation::MISSING) { + if (result) { + Fail("Object(" + object.key + ") exists; expected it not to exist"); + } else if (result.error() != ErrorCode::OBJECT_NOT_FOUND) { + Fail("Object(" + object.key + ") lookup failed with " + + toString(result.error()) + "; expected OBJECT_NOT_FOUND"); + } + return *this; + } if (expectation == ObjectExpectation::NOT_READY) { if (result || result.error() != ErrorCode::REPLICA_IS_NOT_READY) { Fail("Object(" + object.key + ") was expected to be not ready"); @@ -245,6 +293,30 @@ void MasterScenario::ValidateActionResult( } } +void MasterScenario::ValidateStartResult( + std::string_view action, const std::optional& expected_error, + const std::optional& expected_replica_count, + const std::optional& expected_replica_status, + const StartResult& result) { + ValidateActionResult(action, expected_error, result.has_value(), + result ? ErrorCode::OK : result.error()); + if (!result) { + return; + } + if (expected_replica_count.has_value() && + result->size() != *expected_replica_count) { + Fail(std::string(action) + " returned " + + std::to_string(result->size()) + " replicas; expected " + + std::to_string(*expected_replica_count)); + } + if (expected_replica_status.has_value() && + std::any_of(result->begin(), result->end(), [&](const auto& replica) { + return replica.status != *expected_replica_status; + })) { + Fail(std::string(action) + " replica status mismatch"); + } +} + void MasterScenario::Fail(std::string message) const { ADD_FAILURE() << "MasterScenario[" << name_ << "]: " << message; } diff --git a/mooncake-store/tests/master_scenario.h b/mooncake-store/tests/master_scenario.h index 850fa796a9..f6948b7078 100644 --- a/mooncake-store/tests/master_scenario.h +++ b/mooncake-store/tests/master_scenario.h @@ -14,6 +14,8 @@ namespace mooncake::test { +class MasterScenario; + constexpr uint64_t operator""_KB(unsigned long long value) { return value * 1024; } @@ -36,20 +38,33 @@ enum class PutStartExpectation { ERROR, }; +template +struct PutStartAction; + struct PutStartActionData { + PutStartActionData(const PutStartActionData&) = default; + + private: + PutStartActionData(std::string value, uint64_t object_size) + : key(std::move(value)), size(object_size) {} + std::string key; uint64_t size; std::string actor{"default"}; std::optional expected_error{}; std::optional expected_replica_count{}; std::optional expected_replica_status{}; + + template + friend struct PutStartAction; + friend class MasterScenario; }; -template +template struct PutStartAction : PutStartActionData { PutStartAction(std::string key, uint64_t size) requires(expectation == PutStartExpectation::UNSPECIFIED) - : PutStartActionData{.key = std::move(key), .size = size} {} + : PutStartActionData(std::move(key), size) {} PutStartAction& By(std::string value) { actor = std::move(value); @@ -91,6 +106,81 @@ struct PutStartAction : PutStartActionData { PutStartAction<> PutStart(std::string key, uint64_t size); +enum class UpsertStartExpectation { + UNSPECIFIED, + SUCCESS, + ERROR, +}; + +template +struct UpsertStartAction; + +struct UpsertStartActionData { + UpsertStartActionData(const UpsertStartActionData&) = default; + + private: + UpsertStartActionData(std::string value, uint64_t object_size) + : key(std::move(value)), size(object_size) {} + + std::string key; + uint64_t size; + std::string actor{"default"}; + std::optional expected_error{}; + std::optional expected_replica_count{}; + std::optional expected_replica_status{}; + + template + friend struct UpsertStartAction; + friend class MasterScenario; +}; + +template +struct UpsertStartAction : UpsertStartActionData { + UpsertStartAction(std::string key, uint64_t size) + requires(expectation == UpsertStartExpectation::UNSPECIFIED) + : UpsertStartActionData(std::move(key), size) {} + + UpsertStartAction& By(std::string value) { + actor = std::move(value); + return *this; + } + + auto ExpectError(ErrorCode value) const + requires(expectation == UpsertStartExpectation::UNSPECIFIED) + { + UpsertStartAction action(*this); + action.expected_error = value; + return action; + } + + auto ExpectReplicas(size_t value) const + requires(expectation != UpsertStartExpectation::ERROR) + { + UpsertStartAction action(*this); + action.expected_replica_count = value; + return action; + } + + auto ExpectStatus(ReplicaStatus value) const + requires(expectation != UpsertStartExpectation::ERROR) + { + UpsertStartAction action(*this); + action.expected_replica_status = value; + return action; + } + + private: + template + friend struct UpsertStartAction; + + template + UpsertStartAction(const UpsertStartAction& action) + : UpsertStartActionData(action) {} +}; + +UpsertStartAction<> UpsertStart(std::string key, uint64_t size); + struct PutEndAction { std::string key; std::string actor{"default"}; @@ -109,6 +199,24 @@ struct PutEndAction { PutEndAction PutEnd(std::string key); +struct UpsertEndAction { + std::string key; + std::string actor{"default"}; + std::optional expected_error{}; + + UpsertEndAction& By(std::string value) { + actor = std::move(value); + return *this; + } + + UpsertEndAction& ExpectError(ErrorCode value) { + expected_error = value; + return *this; + } +}; + +UpsertEndAction UpsertEnd(std::string key); + struct PutRevokeAction { std::string key; std::string actor{"default"}; @@ -127,6 +235,24 @@ struct PutRevokeAction { PutRevokeAction PutRevoke(std::string key); +struct UpsertRevokeAction { + std::string key; + std::string actor{"default"}; + std::optional expected_error{}; + + UpsertRevokeAction& By(std::string value) { + actor = std::move(value); + return *this; + } + + UpsertRevokeAction& ExpectError(ErrorCode value) { + expected_error = value; + return *this; + } +}; + +UpsertRevokeAction UpsertRevoke(std::string key); + struct RemoveAction { std::string key; std::optional expected_error{}; @@ -143,22 +269,36 @@ enum class ObjectExpectation { UNSPECIFIED, READABLE, NOT_READY, + MISSING, }; +template +struct ObjectSpec; + struct ObjectSpecData { + ObjectSpecData(const ObjectSpecData&) = default; + + private: + explicit ObjectSpecData(std::string value) : key(std::move(value)) {} + std::string key; std::optional expected_replica_count{}; std::optional expected_complete_replica_count{}; + + template + friend struct ObjectSpec; + friend class MasterScenario; }; -template +template struct ObjectSpec : ObjectSpecData { explicit ObjectSpec(std::string key) requires(expectation == ObjectExpectation::UNSPECIFIED) - : ObjectSpecData{.key = std::move(key)} {} + : ObjectSpecData(std::move(key)) {} auto IsReadable() const - requires(expectation != ObjectExpectation::NOT_READY) + requires(expectation != ObjectExpectation::NOT_READY && + expectation != ObjectExpectation::MISSING) { return ObjectSpec(*this); } @@ -169,8 +309,15 @@ struct ObjectSpec : ObjectSpecData { return ObjectSpec(*this); } + auto DoesNotExist() const + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + return ObjectSpec(*this); + } + auto HasReplicas(size_t value) const - requires(expectation != ObjectExpectation::NOT_READY) + requires(expectation != ObjectExpectation::NOT_READY && + expectation != ObjectExpectation::MISSING) { ObjectSpec object(*this); object.expected_replica_count = value; @@ -178,7 +325,8 @@ struct ObjectSpec : ObjectSpecData { } auto HasCompleteReplicas(size_t value) const - requires(expectation != ObjectExpectation::NOT_READY) + requires(expectation != ObjectExpectation::NOT_READY && + expectation != ObjectExpectation::MISSING) { ObjectSpec object(*this); object.expected_complete_replica_count = value; @@ -208,9 +356,15 @@ class MasterScenario { MasterScenario& When(PutStartAction action) { return WhenPutStart(std::move(action)); } + template + MasterScenario& When(UpsertStartAction action) { + return WhenUpsertStart(std::move(action)); + } MasterScenario& When(PutEndAction action); + MasterScenario& When(UpsertEndAction action); MasterScenario& When(PutRevokeAction action); + MasterScenario& When(UpsertRevokeAction action); MasterScenario& When(RemoveAction action); template @@ -220,7 +374,11 @@ class MasterScenario { } private: + using StartResult = + tl::expected, ErrorCode>; + MasterScenario& WhenPutStart(PutStartActionData action); + MasterScenario& WhenUpsertStart(UpsertStartActionData action); MasterScenario& ThenObject(ObjectSpecData object, ObjectExpectation expectation); bool EnsureService(); @@ -228,6 +386,11 @@ class MasterScenario { void ValidateActionResult(std::string_view action, const std::optional& expected_error, bool succeeded, ErrorCode error); + void ValidateStartResult( + std::string_view action, const std::optional& expected_error, + const std::optional& expected_replica_count, + const std::optional& expected_replica_status, + const StartResult& result); void Fail(std::string message) const; std::string name_; diff --git a/mooncake-store/tests/master_scenario_test.cpp b/mooncake-store/tests/master_scenario_test.cpp index bea24d6761..4fc5cb645f 100644 --- a/mooncake-store/tests/master_scenario_test.cpp +++ b/mooncake-store/tests/master_scenario_test.cpp @@ -26,6 +26,25 @@ concept SupportsIsNotReady = requires(T value) { value.IsNotReady(); }; template concept SupportsHasReplicas = requires(T value) { value.HasReplicas(1); }; +template +concept SupportsHasCompleteReplicas = + requires(T value) { value.HasCompleteReplicas(1); }; + +template +concept SupportsDoesNotExist = requires(T value) { value.DoesNotExist(); }; + +template +concept SupportsExpectedErrorMutation = + requires(T value) { value.expected_error = ErrorCode::INTERNAL_ERROR; }; + +template +concept SupportsExpectedReplicaCountMutation = + requires(T value) { value.expected_replica_count = 1; }; + +template +concept SupportsExpectedCompleteReplicaCountMutation = + requires(T value) { value.expected_complete_replica_count = 1; }; + template concept SupportsThen = requires(MasterScenario& scenario, T value) { scenario.Then(value); }; @@ -35,19 +54,41 @@ using ErrorExpectedPutStart = .ExpectError(ErrorCode::INTERNAL_ERROR)); using SuccessExpectedPutStart = decltype(PutStart("compile-time", 1_KB).ExpectReplicas(1)); +using ErrorExpectedUpsertStart = + decltype(UpsertStart("compile-time", 1_KB) + .ExpectError(ErrorCode::INTERNAL_ERROR)); +using SuccessExpectedUpsertStart = + decltype(UpsertStart("compile-time", 1_KB).ExpectReplicas(1)); using UnspecifiedObject = decltype(Object("compile-time")); using NotReadyObject = decltype(Object("compile-time").IsNotReady()); using ReadableObject = decltype(Object("compile-time").HasReplicas(1)); +using MissingObject = decltype(Object("compile-time").DoesNotExist()); static_assert(!SupportsExpectReplicas); static_assert(!SupportsExpectStatus); static_assert(!SupportsExpectError); +static_assert(!SupportsExpectedReplicaCountMutation); +static_assert(!SupportsExpectedErrorMutation); +static_assert(!SupportsExpectReplicas); +static_assert(!SupportsExpectStatus); +static_assert(!SupportsExpectError); +static_assert(!SupportsExpectedReplicaCountMutation); +static_assert(!SupportsExpectedErrorMutation); static_assert(!SupportsIsReadable); static_assert(!SupportsHasReplicas); static_assert(!SupportsIsNotReady); +static_assert(!SupportsDoesNotExist); +static_assert(!SupportsDoesNotExist); +static_assert(!SupportsIsReadable); +static_assert(!SupportsIsNotReady); +static_assert(!SupportsHasReplicas); +static_assert(!SupportsHasCompleteReplicas); +static_assert(!SupportsExpectedReplicaCountMutation); +static_assert(!SupportsExpectedCompleteReplicaCountMutation); static_assert(!SupportsThen); static_assert(SupportsThen); static_assert(SupportsThen); +static_assert(SupportsThen); } // namespace @@ -91,6 +132,58 @@ TEST(MasterScenarioContractTest, ReportsPutStartReplicaStatusMismatch) { "PutStart(key) replica status mismatch"); } +TEST(MasterScenarioContractTest, ReportsUpsertStartReplicaCountMismatch) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("upsert start replica count mismatch") + .Given(MemoryNode("memory")) + .When(UpsertStart("key", 1_KB).ExpectReplicas(2)), + "UpsertStart(key) returned 1 replicas; expected 2"); +} + +TEST(MasterScenarioContractTest, ReportsUpsertStartReplicaStatusMismatch) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("upsert start replica status mismatch") + .Given(MemoryNode("memory")) + .When( + UpsertStart("key", 1_KB).ExpectStatus(ReplicaStatus::COMPLETE)), + "UpsertStart(key) replica status mismatch"); +} + +TEST(MasterScenarioContractTest, ReportsUnexpectedUpsertStartError) { + EXPECT_NONFATAL_FAILURE(MasterScenario("unexpected upsert start error") + .Given(MemoryNode("memory").Capacity(1_KB)) + .When(UpsertStart("key", 2_KB)), + "UpsertStart(key) failed: NO_AVAILABLE_HANDLE"); +} + +TEST(MasterScenarioContractTest, ReportsUnexpectedUpsertStartSuccess) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("unexpected upsert start success") + .Given(MemoryNode("memory")) + .When(UpsertStart("key", 1_KB) + .ExpectError(ErrorCode::OBJECT_ALREADY_EXISTS)), + "UpsertStart(key) succeeded; expected OBJECT_ALREADY_EXISTS"); +} + +TEST(MasterScenarioContractTest, ReportsWrongUpsertEndErrorCode) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("wrong upsert end error code") + .Given(MemoryNode("memory")) + .When(UpsertEnd("missing").ExpectError(ErrorCode::ILLEGAL_CLIENT)), + "UpsertEnd(missing) failed with OBJECT_NOT_FOUND; expected " + "ILLEGAL_CLIENT"); +} + +TEST(MasterScenarioContractTest, ReportsWrongUpsertRevokeErrorCode) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("wrong upsert revoke error code") + .Given(MemoryNode("memory")) + .When( + UpsertRevoke("missing").ExpectError(ErrorCode::ILLEGAL_CLIENT)), + "UpsertRevoke(missing) failed with OBJECT_NOT_FOUND; expected " + "ILLEGAL_CLIENT"); +} + TEST(MasterScenarioContractTest, ReportsUnreadableObject) { EXPECT_NONFATAL_FAILURE( MasterScenario("unreadable object") @@ -99,6 +192,25 @@ TEST(MasterScenarioContractTest, ReportsUnreadableObject) { "Object(missing) is not readable: OBJECT_NOT_FOUND"); } +TEST(MasterScenarioContractTest, ReportsExistingObjectWhenAbsenceExpected) { + EXPECT_NONFATAL_FAILURE(MasterScenario("existing object expected absent") + .Given(MemoryNode("memory")) + .When(PutStart("key", 1_KB)) + .When(PutEnd("key")) + .Then(Object("key").DoesNotExist()), + "Object(key) exists; expected it not to exist"); +} + +TEST(MasterScenarioContractTest, ReportsNotReadyObjectWhenAbsenceExpected) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("not-ready object expected absent") + .Given(MemoryNode("memory")) + .When(PutStart("key", 1_KB)) + .Then(Object("key").DoesNotExist()), + "Object(key) lookup failed with REPLICA_IS_NOT_READY; expected " + "OBJECT_NOT_FOUND"); +} + TEST(MasterScenarioContractTest, ReportsObjectReplicaCountMismatch) { EXPECT_NONFATAL_FAILURE(MasterScenario("object replica count mismatch") .Given(MemoryNode("memory")) diff --git a/mooncake-store/tests/master_service_scenario_test.cpp b/mooncake-store/tests/master_service_scenario_test.cpp index ba342bf151..822925ec14 100644 --- a/mooncake-store/tests/master_service_scenario_test.cpp +++ b/mooncake-store/tests/master_service_scenario_test.cpp @@ -26,4 +26,113 @@ TEST(MasterServiceTest, PutStartEndFlow) { .HasCompleteReplicas(1)); } +TEST(MasterServiceTest, UpsertNewKey) { + MasterScenario("upsert creates a new object") + .Given(MemoryNode("memory")) + .When(UpsertStart("upsert_new_key", 1_KB) + .By("writer") + .ExpectReplicas(1) + .ExpectStatus(ReplicaStatus::PROCESSING)) + .Then(Object("upsert_new_key").IsNotReady()) + .When(UpsertEnd("upsert_new_key").By("writer")) + .Then(Object("upsert_new_key") + .IsReadable() + .HasReplicas(1) + .HasCompleteReplicas(1)); +} + +TEST(MasterServiceTest, UpsertPreemptsInProgressPut) { + MasterScenario("upsert preempts an in-progress put") + .Given(MemoryNode("memory")) + .When(PutStart("upsert_preempt", 1_KB).By("old-writer")) + .When(UpsertStart("upsert_preempt", 1_KB) + .By("new-writer") + .ExpectReplicas(1) + .ExpectStatus(ReplicaStatus::PROCESSING)) + .When(PutEnd("upsert_preempt") + .By("old-writer") + .ExpectError(ErrorCode::ILLEGAL_CLIENT)) + .When(UpsertEnd("upsert_preempt").By("new-writer")) + .Then(Object("upsert_preempt") + .IsReadable() + .HasReplicas(1) + .HasCompleteReplicas(1)); +} + +TEST(MasterServiceTest, UpsertSameSizeRefreshesMetadata) { + MasterScenario("same-size upsert refreshes writer identity") + .Given(MemoryNode("memory")) + .When(PutStart("upsert_refresh_metadata", 1_KB).By("old-writer")) + .When(PutEnd("upsert_refresh_metadata").By("old-writer")) + .When(UpsertStart("upsert_refresh_metadata", 1_KB).By("new-writer")) + .When(UpsertEnd("upsert_refresh_metadata") + .By("old-writer") + .ExpectError(ErrorCode::ILLEGAL_CLIENT)) + .When(UpsertEnd("upsert_refresh_metadata").By("new-writer")) + .Then(Object("upsert_refresh_metadata") + .IsReadable() + .HasReplicas(1) + .HasCompleteReplicas(1)); +} + +TEST(MasterServiceTest, UpsertRevoke) { + MasterScenario("revoke a new-key upsert") + .Given(MemoryNode("memory")) + .When(UpsertStart("upsert_revoke", 1_KB).By("writer")) + .Then(Object("upsert_revoke").IsNotReady()) + .When(UpsertRevoke("upsert_revoke").By("writer")) + .Then(Object("upsert_revoke").DoesNotExist()); +} + +TEST(MasterServiceTest, UpsertInPlaceThenRevoke) { + MasterScenario("revoke an in-place upsert") + .Given(MemoryNode("memory")) + .When(PutStart("upsert_inplace_revoke", 1_KB).By("old-writer")) + .When(PutEnd("upsert_inplace_revoke").By("old-writer")) + .When(UpsertStart("upsert_inplace_revoke", 1_KB).By("new-writer")) + .Then(Object("upsert_inplace_revoke").IsNotReady()) + .When(UpsertRevoke("upsert_inplace_revoke").By("new-writer")) + .Then(Object("upsert_inplace_revoke").DoesNotExist()); +} + +TEST(MasterServiceTest, UpsertPreemptsInProgressUpsert) { + MasterScenario("upsert preempts an in-progress upsert") + .Given(MemoryNode("memory")) + .When(PutStart("upsert_preempt_upsert", 1_KB).By("first-writer")) + .When(PutEnd("upsert_preempt_upsert").By("first-writer")) + .When( + UpsertStart("upsert_preempt_upsert", 1_KB).By("old-upsert-writer")) + .Then(Object("upsert_preempt_upsert").IsNotReady()) + .When(UpsertStart("upsert_preempt_upsert", 1_KB) + .By("new-upsert-writer") + .ExpectReplicas(1) + .ExpectStatus(ReplicaStatus::PROCESSING)) + .When(UpsertEnd("upsert_preempt_upsert") + .By("old-upsert-writer") + .ExpectError(ErrorCode::ILLEGAL_CLIENT)) + .When(UpsertEnd("upsert_preempt_upsert").By("new-upsert-writer")) + .Then(Object("upsert_preempt_upsert") + .IsReadable() + .HasReplicas(1) + .HasCompleteReplicas(1)); +} + +TEST(MasterServiceTest, UpsertDifferentSizeThenRevoke) { + MasterScenario("revoke a different-size upsert") + .Given(MemoryNode("memory")) + .When(PutStart("upsert_diff_revoke", 1_KB).By("writer")) + .When(PutEnd("upsert_diff_revoke").By("writer")) + .Then(Object("upsert_diff_revoke") + .IsReadable() + .HasReplicas(1) + .HasCompleteReplicas(1)) + .When(UpsertStart("upsert_diff_revoke", 2_KB) + .By("writer") + .ExpectReplicas(1) + .ExpectStatus(ReplicaStatus::PROCESSING)) + .Then(Object("upsert_diff_revoke").IsNotReady()) + .When(UpsertRevoke("upsert_diff_revoke").By("writer")) + .Then(Object("upsert_diff_revoke").DoesNotExist()); +} + } // namespace mooncake::test diff --git a/mooncake-store/tests/master_service_test.cpp b/mooncake-store/tests/master_service_test.cpp index f8c453cbdc..9068dc156f 100644 --- a/mooncake-store/tests/master_service_test.cpp +++ b/mooncake-store/tests/master_service_test.cpp @@ -6809,41 +6809,6 @@ TEST_F(MasterServiceTest, ForceRemoveAllLeasedObjects) { } // ===================== Upsert Tests ===================== -TEST_F(MasterServiceTest, UpsertNewKey) { - // Case A: key does not exist — behaves like PutStart - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - - std::string key = "upsert_new_key"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - auto upsert_result = service_->UpsertStart( - client_id, key, TenantId::Default(), slice_length, config); - ASSERT_TRUE(upsert_result.has_value()); - auto replicas = upsert_result.value(); - EXPECT_EQ(1, replicas.size()); - EXPECT_EQ(ReplicaStatus::PROCESSING, replicas[0].status); - - // During upsert, GetReplicaList should return not ready - auto get_result = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_FALSE(get_result.has_value()); - EXPECT_EQ(ErrorCode::REPLICA_IS_NOT_READY, get_result.error()); - - // UpsertEnd completes the operation - auto end_result = service_->UpsertEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(end_result.has_value()); - - // Verify replica is COMPLETE - auto final_result = service_->GetReplicaList(key, TenantId::Default()); - ASSERT_TRUE(final_result.has_value()); - EXPECT_EQ(1, final_result.value().replicas.size()); - EXPECT_EQ(ReplicaStatus::COMPLETE, final_result.value().replicas[0].status); -} - TEST_F(MasterServiceTest, UpsertSameSize) { // Case B: key exists with same size — in-place update std::unique_ptr service_(new MasterService()); @@ -6892,43 +6857,6 @@ TEST_F(MasterServiceTest, UpsertSameSize) { EXPECT_EQ(ReplicaStatus::COMPLETE, final_result.value().replicas[0].status); } -TEST_F(MasterServiceTest, UpsertSameSizeRefreshesMetadata) { - // Case B: verify client_id and put_start_time are refreshed - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id_a = generate_uuid(); - const UUID client_id_b = generate_uuid(); - - std::string key = "upsert_refresh_metadata"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - // Create object with client_a - auto put_result = service_->PutStart(client_id_a, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(put_result.has_value()); - auto put_end = service_->PutEnd(client_id_a, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end.has_value()); - - // UpsertStart with client_b - auto upsert_result = service_->UpsertStart( - client_id_b, key, TenantId::Default(), slice_length, config); - ASSERT_TRUE(upsert_result.has_value()); - - // UpsertEnd with client_a should fail (client_id was refreshed to client_b) - auto end_fail = service_->UpsertEnd(client_id_a, key, TenantId::Default(), - ReplicaType::MEMORY); - EXPECT_FALSE(end_fail.has_value()); - EXPECT_EQ(ErrorCode::ILLEGAL_CLIENT, end_fail.error()); - - // UpsertEnd with client_b should succeed - auto end_ok = service_->UpsertEnd(client_id_b, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(end_ok.has_value()); -} - TEST_F(MasterServiceTest, UpsertDifferentSize) { // Case C: key exists with different size — delete and reallocate std::unique_ptr service_(new MasterService()); @@ -7017,111 +6945,6 @@ TEST_F(MasterServiceTest, UpsertConflictReplicationTask) { EXPECT_EQ(ErrorCode::OBJECT_HAS_REPLICATION_TASK, upsert_result.error()); } -TEST_F(MasterServiceTest, UpsertPreemptsInProgressPut) { - // Upsert should preempt an in-progress Put (no discard timeout needed - // for preemption via Upsert — Upsert always preempts immediately) - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_a = generate_uuid(); - const UUID client_b = generate_uuid(); - - std::string key = "upsert_preempt"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - // Client A starts a Put but doesn't finish - auto put_result = service_->PutStart(client_a, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(put_result.has_value()); - - // Client B upserts the same key — should preempt client A - auto upsert_result = service_->UpsertStart( - client_b, key, TenantId::Default(), slice_length, config); - ASSERT_TRUE(upsert_result.has_value()); - auto upsert_replicas = upsert_result.value(); - EXPECT_EQ(1, upsert_replicas.size()); - EXPECT_EQ(ReplicaStatus::PROCESSING, upsert_replicas[0].status); - - // Client A's PutEnd should fail - auto put_end_a = service_->PutEnd(client_a, key, TenantId::Default(), - ReplicaType::MEMORY); - EXPECT_FALSE(put_end_a.has_value()); - - // Client B's UpsertEnd should succeed - auto upsert_end = service_->UpsertEnd(client_b, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(upsert_end.has_value()); - - // Verify final state - auto final_result = service_->GetReplicaList(key, TenantId::Default()); - ASSERT_TRUE(final_result.has_value()); - EXPECT_EQ(ReplicaStatus::COMPLETE, final_result.value().replicas[0].status); -} - -TEST_F(MasterServiceTest, UpsertRevoke) { - // UpsertRevoke should clean up like PutRevoke - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - - std::string key = "upsert_revoke"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - // UpsertStart (Case A — new key) - auto upsert_result = service_->UpsertStart( - client_id, key, TenantId::Default(), slice_length, config); - ASSERT_TRUE(upsert_result.has_value()); - - // UpsertRevoke - auto revoke_result = service_->UpsertRevoke( - client_id, key, TenantId::Default(), ReplicaType::MEMORY); - ASSERT_TRUE(revoke_result.has_value()); - - // Key should be gone - auto exist_result = service_->ExistKey(key, TenantId::Default()); - ASSERT_TRUE(exist_result.has_value()); - EXPECT_FALSE(exist_result.value()); -} - -TEST_F(MasterServiceTest, UpsertInPlaceThenRevoke) { - // UpsertRevoke after in-place UpsertStart should clean up - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - - std::string key = "upsert_inplace_revoke"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - // Create object first - auto put_result = service_->PutStart(client_id, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(put_result.has_value()); - auto put_end = service_->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end.has_value()); - - // UpsertStart in-place (same size) - const UUID new_client = generate_uuid(); - auto upsert_result = service_->UpsertStart( - new_client, key, TenantId::Default(), slice_length, config); - ASSERT_TRUE(upsert_result.has_value()); - - // UpsertRevoke — replicas are PROCESSING, should be erased - auto revoke_result = service_->UpsertRevoke( - new_client, key, TenantId::Default(), ReplicaType::MEMORY); - ASSERT_TRUE(revoke_result.has_value()); - - // Key should be gone (no valid replicas left) - auto exist_result = service_->ExistKey(key, TenantId::Default()); - ASSERT_TRUE(exist_result.has_value()); - EXPECT_FALSE(exist_result.value()); -} - TEST_F(MasterServiceTest, BatchUpsertStart) { // Test batch upsert with a mix of new and existing keys std::unique_ptr service_(new MasterService()); @@ -7157,107 +6980,6 @@ TEST_F(MasterServiceTest, BatchUpsertStart) { EXPECT_TRUE(end_results[1].has_value()); } -TEST_F(MasterServiceTest, UpsertPreemptsInProgressUpsert) { - // Upsert should preempt an in-progress Upsert (Case B in-place). - // After preemption, all replicas were PROCESSING (no COMPLETE survives), - // so metadata is erased and the new upsert falls through to Case A. - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_a = generate_uuid(); - const UUID client_b = generate_uuid(); - const UUID client_c = generate_uuid(); - - std::string key = "upsert_preempt_upsert"; - uint64_t slice_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - - // Step 1: Create the object via Put - auto put_result = service_->PutStart(client_a, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(put_result.has_value()); - auto put_end = service_->PutEnd(client_a, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end.has_value()); - - // Step 2: Client B starts in-place upsert (Case B) — marks COMPLETE → - // PROCESSING - auto upsert_b = service_->UpsertStart(client_b, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(upsert_b.has_value()); - - // Key should be unreadable now (all replicas are PROCESSING) - auto get_mid = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_FALSE(get_mid.has_value()); - EXPECT_EQ(ErrorCode::REPLICA_IS_NOT_READY, get_mid.error()); - - // Step 3: Client C upserts the same key — preempts Client B - auto upsert_c = service_->UpsertStart(client_c, key, TenantId::Default(), - slice_length, config); - ASSERT_TRUE(upsert_c.has_value()); - EXPECT_EQ(1, upsert_c.value().size()); - - // Step 4: Client B's UpsertEnd should fail (preempted) - auto end_b = service_->UpsertEnd(client_b, key, TenantId::Default(), - ReplicaType::MEMORY); - EXPECT_FALSE(end_b.has_value()); - - // Step 5: Client C's UpsertEnd should succeed - auto end_c = service_->UpsertEnd(client_c, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(end_c.has_value()); - - // Final verification - auto final_result = service_->GetReplicaList(key, TenantId::Default()); - ASSERT_TRUE(final_result.has_value()); - EXPECT_EQ(1, final_result.value().replicas.size()); - EXPECT_EQ(ReplicaStatus::COMPLETE, final_result.value().replicas[0].status); -} - -TEST_F(MasterServiceTest, UpsertDifferentSizeThenRevoke) { - // Case C (different size) followed by UpsertRevoke. - // Old replicas go to discarded_replicas_, new replicas are erased by - // revoke. The key should disappear entirely. - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - - std::string key = "upsert_diff_revoke"; - uint64_t original_size = 1024; - uint64_t new_size = 2048; - ReplicateConfig config; - config.replica_num = 1; - - // Create object with original size - auto put_result = service_->PutStart(client_id, key, TenantId::Default(), - original_size, config); - ASSERT_TRUE(put_result.has_value()); - auto put_end = service_->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end.has_value()); - - // Verify the key exists - auto exist_before = service_->ExistKey(key, TenantId::Default()); - ASSERT_TRUE(exist_before.has_value()); - EXPECT_TRUE(exist_before.value()); - - // UpsertStart with different size (Case C) — old replicas discarded, - // new replicas allocated - auto upsert_result = service_->UpsertStart( - client_id, key, TenantId::Default(), new_size, config); - ASSERT_TRUE(upsert_result.has_value()); - - // Revoke — erase the newly allocated PROCESSING replicas - auto revoke_result = service_->UpsertRevoke( - client_id, key, TenantId::Default(), ReplicaType::MEMORY); - ASSERT_TRUE(revoke_result.has_value()); - - // Key should be gone (old replicas in discarded, new replicas erased) - auto exist_after = service_->ExistKey(key, TenantId::Default()); - ASSERT_TRUE(exist_after.has_value()); - EXPECT_FALSE(exist_after.value()); -} - // ===================== Hard Pin Tests ===================== TEST_F(MasterServiceTest, HardPinObjectNotEvicted) { From a7352a2f5fb96a919481fbe953c6812c2f734119 Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Mon, 10 Aug 2026 14:17:33 +0800 Subject: [PATCH 014/483] [Wheel] add DataProto catalog for fragmented rollout data (#3345) --- mooncake-wheel/mooncake/dataproto_catalog.py | 238 ++++++++++++++++++ .../mooncake/structured_object_store.py | 39 ++- .../tests/test_dataproto_catalog.py | 122 +++++++++ .../tests/test_structured_object_store.py | 41 +++ 4 files changed, 437 insertions(+), 3 deletions(-) create mode 100644 mooncake-wheel/mooncake/dataproto_catalog.py create mode 100644 mooncake-wheel/tests/test_dataproto_catalog.py diff --git a/mooncake-wheel/mooncake/dataproto_catalog.py b/mooncake-wheel/mooncake/dataproto_catalog.py new file mode 100644 index 0000000000..f30d9f48da --- /dev/null +++ b/mooncake-wheel/mooncake/dataproto_catalog.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +import copy +from collections.abc import Mapping, Sequence +from typing import Any + + +class DataProtoCatalog: + """Map logical rows and fields to immutable DataProto refs. + + The catalog only tracks metadata. Data transfer and object lifetime remain + owned by MooncakeBundleTransfer and Mooncake Store, respectively. Calls to + one catalog instance must be serialized by its host. Each update publishes + one immutable, single-stage fragment; multi-stage append handles are not + catalog fragments. + """ + + def __init__(self) -> None: + self._partitions: dict[str, dict[str, dict[str, Any]]] = {} + self._fragments: dict[str, dict[str, Any]] = {} + self._fragment_refcounts: dict[str, int] = {} + + def update( + self, + partition: str, + keys: Sequence[str], + *, + tags: Sequence[Mapping[str, Any]] | None = None, + handle: Mapping[str, Any] | None = None, + ) -> dict[str, Any]: + """Merge tags and publish every field in a single-stage ``handle``.""" + partition = _nonempty_string(partition, "partition") + keys = _names(keys, "keys") + normalized_tags = _tags(tags, len(keys)) + fragment = _fragment(handle, len(keys)) if handle is not None else None + if normalized_tags is None and fragment is None: + raise ValueError("catalog update requires tags or a DataProto handle") + + if fragment is not None: + fragment_id, stored_handle, fields = fragment + previous = self._fragments.get(fragment_id) + if previous is not None and previous != stored_handle: + raise ValueError(f"DataProto fragment id collision: {fragment_id!r}") + + current_entries = self._partitions.get(partition, {}) + merged_tags = {} + response_tags = [] + for row, key in enumerate(keys): + tag = copy.deepcopy(current_entries.get(key, {"tag": {}})["tag"]) + if normalized_tags is not None: + tag.update(normalized_tags[row]) + merged_tags[key] = tag + response_tags.append(copy.deepcopy(tag)) + + entries = self._partitions.setdefault(partition, {}) + if fragment is not None: + self._fragments.setdefault(fragment_id, stored_handle) + self._fragment_refcounts.setdefault(fragment_id, 0) + + for row, key in enumerate(keys): + entry = entries.setdefault(key, {"tag": {}, "fields": {}}) + entry["tag"] = merged_tags[key] + if fragment is None: + continue + for field in fields: + location = (fragment_id, row) + old_location = entry["fields"].get(field) + if old_location == location: + continue + if old_location is not None and old_location[0] == fragment_id: + entry["fields"][field] = location + continue + if old_location is not None: + self._drop_fragment_reference(old_location[0]) + entry["fields"][field] = location + self._fragment_refcounts[fragment_id] += 1 + + return { + "keys": list(keys), + "tags": response_tags, + "fields": _field_union(entries, keys), + } + + def resolve( + self, + partition: str, + keys: Sequence[str], + fields: Sequence[str] | None = None, + ) -> dict[str, Any]: + """Resolve an ordered logical read into immutable fragment locations.""" + partition = _nonempty_string(partition, "partition") + keys = _names(keys, "keys") + entries = self._partitions.get(partition, {}) + missing = [key for key in keys if key not in entries] + if missing: + raise ValueError( + f"keys were not found in partition {partition!r}: {missing}" + ) + + selected = ( + list(_names(fields, "field names")) + if fields is not None + else _field_union(entries, keys) + ) + if not selected: + raise ValueError("requested keys do not contain any fields") + + for key in keys: + entry_fields = entries[key]["fields"] + unavailable = [field for field in selected if field not in entry_fields] + if unavailable: + raise ValueError(f"fields are not ready for key {key!r}: {unavailable}") + + grouped_fields: dict[tuple[tuple[str, int], ...], list[str]] = {} + fragment_ids: set[str] = set() + for field in selected: + locations = tuple(entries[key]["fields"][field] for key in keys) + grouped_fields.setdefault(locations, []).append(field) + fragment_ids.update(location[0] for location in locations) + + return { + "keys": list(keys), + "tags": [copy.deepcopy(entries[key]["tag"]) for key in keys], + "fields": selected, + "field_groups": [ + {"fields": fields, "locations": list(locations)} + for locations, fields in grouped_fields.items() + ], + "handles": { + fragment_id: copy.deepcopy(self._fragments[fragment_id]) + for fragment_id in fragment_ids + }, + } + + def list(self, partition: str | None = None) -> dict[str, dict[str, Any]]: + """Return tags using the same partition/key shape as rollout KV APIs.""" + if partition is not None: + partition = _nonempty_string(partition, "partition") + partitions = [partition] if partition in self._partitions else [] + else: + partitions = list(self._partitions) + return { + partition_id: { + key: copy.deepcopy(entry["tag"]) + for key, entry in self._partitions[partition_id].items() + } + for partition_id in partitions + } + + def remove(self, partition: str, keys: Sequence[str]) -> None: + """Remove logical keys after their writers have finished. + + This only removes catalog metadata; physical lifetime stays with Store. + """ + partition = _nonempty_string(partition, "partition") + keys = _names(keys, "keys") + entries = self._partitions.get(partition) + if entries is None: + return + + for key in keys: + entry = entries.pop(key, None) + if entry is None: + continue + for fragment_id, _row in entry["fields"].values(): + self._drop_fragment_reference(fragment_id) + if not entries: + del self._partitions[partition] + + def _drop_fragment_reference(self, fragment_id: str) -> None: + remaining = self._fragment_refcounts[fragment_id] - 1 + if remaining: + self._fragment_refcounts[fragment_id] = remaining + return + del self._fragment_refcounts[fragment_id] + del self._fragments[fragment_id] + + +def _fragment( + handle: Mapping[str, Any], batch_size: int +) -> tuple[str, dict[str, Any], tuple[str, ...]]: + from mooncake.structured_object_store import ( + export_dataproto_ref, + import_dataproto_ref, + ) + + ref = import_dataproto_ref(handle) + if ref.batch_size != batch_size: + raise ValueError( + f"DataProto batch size {ref.batch_size!r} does not match " + f"{batch_size} logical keys" + ) + if len(ref.stage_refs) != 1: + raise ValueError("catalog fragments must contain exactly one DataProto stage") + fragment_id = _nonempty_string( + next(iter(ref.stage_refs.values())).manifest_key, "manifest_key" + ) + fields = _names(ref.field_index, "field names") + if not fields: + raise ValueError("DataProto catalog fragments cannot be empty") + return fragment_id, export_dataproto_ref(ref), fields + + +def _field_union( + entries: Mapping[str, Mapping[str, Any]], keys: Sequence[str] +) -> list[str]: + return list( + dict.fromkeys(field for key in keys for field in entries[key]["fields"]) + ) + + +def _names(values: Sequence[str] | Mapping[str, Any], label: str) -> tuple[str, ...]: + if isinstance(values, str): + raise ValueError(f"{label} must be a sequence of non-empty strings") + result = tuple(values) + if not result or any(not isinstance(value, str) or not value for value in result): + raise ValueError(f"{label} must be non-empty strings") + if len(result) != len(set(result)): + raise ValueError(f"{label} must be unique") + return result + + +def _tags( + tags: Sequence[Mapping[str, Any]] | None, count: int +) -> list[dict[str, Any]] | None: + if tags is None: + return None + if len(tags) != count: + raise ValueError("tags must have the same length as keys") + if any(not isinstance(tag, Mapping) for tag in tags): + raise TypeError("tags must be mappings") + return [copy.deepcopy(dict(tag)) for tag in tags] + + +def _nonempty_string(value: Any, name: str) -> str: + if not isinstance(value, str) or not value: + raise ValueError(f"{name} must be a non-empty string") + return value diff --git a/mooncake-wheel/mooncake/structured_object_store.py b/mooncake-wheel/mooncake/structured_object_store.py index 6eee7df21e..1985ec00e3 100644 --- a/mooncake-wheel/mooncake/structured_object_store.py +++ b/mooncake-wheel/mooncake/structured_object_store.py @@ -363,9 +363,14 @@ def import_dataproto_ref(handle: Mapping[str, Any]) -> MooncakeDataProtoRef: manifest_key = _require_mapping(stage_ref, f"stage_refs[{stage!r}]").get( "manifest_key" ) - if not isinstance(stage, str) or not isinstance(manifest_key, str): + if ( + not isinstance(stage, str) + or not stage + or not isinstance(manifest_key, str) + or not manifest_key + ): raise ValueError( - "DataProto ref handle stage refs must contain string manifest_key values" + "DataProto ref handle stage refs must contain non-empty string names and manifest_key values" ) stage_refs[stage] = RemoteBundleRef(manifest_key=manifest_key, manifest={}) field_index: dict[str, StructuredFieldLocation] = {} @@ -380,10 +385,15 @@ def import_dataproto_ref(handle: Mapping[str, Any]) -> MooncakeDataProtoRef: member = location.get("member") if ( not isinstance(name, str) + or not name or not isinstance(stage, str) + or not stage or not isinstance(member, str) + or not member ): - raise ValueError("DataProto ref handle field locations must be strings") + raise ValueError( + "DataProto ref handle field locations must be non-empty strings" + ) if stage not in stage_refs: raise ValueError( f"DataProto ref handle field {name!r} references unknown stage {stage!r}" @@ -999,6 +1009,15 @@ def _read_torch_tensor_member_indices( indices: Sequence[int], destination: Any, ) -> Any: + if payload_spec.get("format") == "torch_save": + if destination is not None: + raise ValueError( + f"structured torch_save tensor member {name} does not support destinations" + ) + value = _deserialize_torch_save_payload( + self._bundle_store.read_payload(payload_spec) + ) + return value[list(indices)] metadata_bytes = int(payload_spec.get("metadata_bytes", -1)) shape = field_spec.get("shape") element_size = int(field_spec.get("element_size", 0)) @@ -3301,6 +3320,20 @@ def _read_sliced_torch_tensor_member( member_slice: StructuredMemberSlice, destination: Any, ) -> Any: + if payload_spec.get("format") == "torch_save": + if destination is not None: + raise ValueError( + f"structured torch_save tensor member {name} does not support destinations" + ) + if member_slice.axis != 0: + raise ValueError( + "structured tensor slicing currently supports axis=0 only" + ) + value = _deserialize_torch_save_payload( + self._bundle_store.read_payload(payload_spec) + ) + start, end, step = _normalized_member_slice(member_slice, len(value)) + return value[start:end:step] metadata_bytes = int(payload_spec.get("metadata_bytes", -1)) shape = field_spec.get("shape") element_size = int(field_spec.get("element_size", 0)) diff --git a/mooncake-wheel/tests/test_dataproto_catalog.py b/mooncake-wheel/tests/test_dataproto_catalog.py new file mode 100644 index 0000000000..ac81c420cf --- /dev/null +++ b/mooncake-wheel/tests/test_dataproto_catalog.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +import pytest + +from mooncake.dataproto_catalog import DataProtoCatalog + + +def handle(name: str, fields: list[str], batch_size: int = 2) -> dict: + return { + "type": "mooncake_dataproto_ref", + "version": 1, + "kind": "bundle_stages", + "batch_size": batch_size, + "stage_refs": {"rollout": {"manifest_key": f"manifest/{name}"}}, + "field_index": { + field: { + "stage": "rollout", + "member": f"batch.{field}", + "section": "batch", + } + for field in fields + }, + } + + +def test_catalog_merges_tags_and_resolves_fragmented_fields_in_key_order() -> None: + catalog = DataProtoCatalog() + catalog.update( + "train", + ["a", "b"], + tags=[{"status": "running"}, {"status": "running"}], + handle=handle("base", ["input_ids", "attention_mask"]), + ) + result = catalog.update( + "train", + ["b", "a"], + tags=[{"status": "done"}, {"status": "done"}], + handle=handle("scores", ["score"]), + ) + + assert result["fields"] == ["input_ids", "attention_mask", "score"] + assert catalog.list("train") == { + "train": { + "a": {"status": "done"}, + "b": {"status": "done"}, + } + } + plan = catalog.resolve("train", ["a", "b"], ["score", "input_ids"]) + assert plan["fields"] == ["score", "input_ids"] + assert plan["field_groups"] == [ + { + "fields": ["score"], + "locations": [("manifest/scores", 1), ("manifest/scores", 0)], + }, + { + "fields": ["input_ids"], + "locations": [("manifest/base", 0), ("manifest/base", 1)], + }, + ] + assert set(plan["handles"]) == {"manifest/base", "manifest/scores"} + grouped = catalog.resolve("train", ["b", "a"], ["input_ids", "attention_mask"]) + assert catalog.resolve("train", ["a", "b"])["fields"] == [ + "input_ids", + "attention_mask", + "score", + ] + assert grouped["field_groups"] == [ + { + "fields": ["input_ids", "attention_mask"], + "locations": [("manifest/base", 1), ("manifest/base", 0)], + } + ] + + +def test_catalog_drops_replaced_and_removed_locations() -> None: + catalog = DataProtoCatalog() + base = handle("base", ["value"]) + replacement = handle("replacement", ["value"]) + catalog.update("train", ["a", "b"], handle=base) + + catalog.update("train", ["a"], handle={**replacement, "batch_size": 1}) + catalog.update("train", ["b", "a"], handle=base) + plan = catalog.resolve("train", ["a", "b"], ["value"]) + assert plan["field_groups"] == [ + { + "fields": ["value"], + "locations": [("manifest/base", 1), ("manifest/base", 0)], + } + ] + catalog.remove("train", ["b"]) + catalog.remove("train", ["a"]) + assert catalog.list() == {} + + +def test_catalog_rejects_missing_or_incomplete_reads_without_mutation() -> None: + catalog = DataProtoCatalog() + catalog.update("train", ["a", "b"], tags=[{}, {}]) + catalog.update("train", ["a"], handle=handle("one", ["value"], batch_size=1)) + + with pytest.raises(ValueError, match="were not found"): + catalog.resolve("train", ["missing"]) + with pytest.raises(ValueError, match="not ready"): + catalog.resolve("train", ["a", "b"], ["value"]) + with pytest.raises(ValueError, match="do not contain any fields"): + catalog.resolve("train", ["b"]) + with pytest.raises(ValueError, match="same length"): + catalog.update("train", ["a", "b"], tags=[{}]) + with pytest.raises(ValueError, match="sequence"): + catalog.resolve("train", "a") + with pytest.raises(ValueError, match="sequence"): + catalog.resolve("train", ["a"], "value") + invalid_handle = handle("invalid", ["value"], batch_size=1) + invalid_handle["version"] = 999 + with pytest.raises(ValueError, match="unsupported.*version"): + catalog.update("invalid", ["a"], handle=invalid_handle) + multi_stage = handle("multi", ["value"], batch_size=1) + multi_stage["stage_refs"]["extra"] = {"manifest_key": "manifest/extra"} + with pytest.raises(ValueError, match="exactly one.*stage"): + catalog.update("invalid", ["a"], handle=multi_stage) + + assert catalog.list("train") == {"train": {"a": {}, "b": {}}} + assert "invalid" not in catalog.list() diff --git a/mooncake-wheel/tests/test_structured_object_store.py b/mooncake-wheel/tests/test_structured_object_store.py index 971c4d89b5..d65a831f4a 100644 --- a/mooncake-wheel/tests/test_structured_object_store.py +++ b/mooncake-wheel/tests/test_structured_object_store.py @@ -526,6 +526,27 @@ def test_put_object_torch_tensor_raw_fallback_roundtrip() -> None: assert store.batch_get_into_calls > 0 +def test_dataproto_torch_save_fallback_supports_row_selection(monkeypatch) -> None: + torch = pytest.importorskip("torch") + monkeypatch.setattr(sos, "_has_tensor_codec_helpers", lambda: False) + _store, transfer = make_transfer(NoTensorFastPathStore()) + tensor = torch.arange(24, dtype=torch.float32).reshape(6, 4) + + ref = transfer.put_dataproto(SimpleDataProto(batch={"tensor": tensor})) + payload = ref.stage_refs["default"].manifest["buffers"]["batch.tensor"] + + assert payload["format"] == "torch_save" + for rows in (slice(1, 6, 2), [4, 1, 4], [], [-1, 0]): + assert torch.equal( + transfer.get_dataproto(ref, rows=rows)["batch"]["tensor"], + tensor[rows], + ) + with pytest.raises(ValueError, match="does not support destinations"): + transfer.get_dataproto( + ref, rows=[1], destinations={"tensor": object()} + ) + + def test_bundle_read_spec_full_read_is_partial_special_case() -> None: store, transfer = make_transfer() array = np.arange(16, dtype=np.int32).reshape(4, 4) @@ -2110,6 +2131,26 @@ def test_dataproto_ref_handle_rejects_unknown_field_stage() -> None: import_dataproto_ref(handle) +@pytest.mark.parametrize( + ("path", "value"), + [ + (("stage_refs", "default", "manifest_key"), ""), + (("field_index", "input_ids", "member"), ""), + ], +) +def test_dataproto_ref_handle_rejects_empty_locations(path, value) -> None: + _store, transfer = make_transfer() + ref = transfer.put_dataproto(SimpleDataProto(batch={"input_ids": np.arange(4)})) + handle = export_dataproto_ref(ref) + target = handle + for key in path[:-1]: + target = target[key] + target[path[-1]] = value + + with pytest.raises(ValueError, match="non-empty"): + import_dataproto_ref(handle) + + def test_dataproto_helper_appends_stage_fields_without_rewriting_existing() -> None: store, transfer = make_transfer() rollout = SimpleDataProto( From 78726c5ca6417f3926c384320728bb7c426e8f22 Mon Sep 17 00:00:00 2001 From: Xiao You Date: Mon, 10 Aug 2026 14:19:44 +0800 Subject: [PATCH 015/483] [Store] Add put/get session APIs for ranged multi-buffer transfers (#2881) Introduce session-scoped put/get APIs for layer-wise ranged multi-buffer transfers, avoiding repeated Master RPCs between layers. Address review feedback on session put finalize: - Finalize MEMORY replicas only and revoke leftover NoF reservations, since BatchTransferWriteRanges does not write NoF. - Keep put sessions on BatchPutEnd/BatchPutRevoke RPC ambiguity or per-key failure so callers can retry end or revoke. Co-authored-by: youyx Co-authored-by: Cursor --- .../api-reference/python/mooncake-store.md | 191 ++++++ mooncake-integration/store/store_py.cpp | 105 +++ mooncake-store/include/client_service.h | 58 ++ mooncake-store/include/pyclient.h | 51 ++ mooncake-store/include/real_client.h | 46 ++ mooncake-store/include/transfer_task.h | 8 + mooncake-store/src/client_service.cpp | 307 ++++++++- mooncake-store/src/real_client.cpp | 620 +++++++++++++++++- mooncake-store/src/transfer_task.cpp | 58 +- mooncake-store/tests/client_buffer_test.cpp | 6 +- .../tests/e2e/run_session_ranges_tcp_e2e.sh | 79 +++ .../tests/e2e/session_ranges_tcp_e2e.py | 152 +++++ mooncake-store/tests/pybind_client_test.cpp | 480 ++++++++++++++ 13 files changed, 2152 insertions(+), 9 deletions(-) create mode 100755 mooncake-store/tests/e2e/run_session_ranges_tcp_e2e.sh create mode 100644 mooncake-store/tests/e2e/session_ranges_tcp_e2e.py diff --git a/docs/source/api-reference/python/mooncake-store.md b/docs/source/api-reference/python/mooncake-store.md index fe6d71d3da..dd90000cbe 100644 --- a/docs/source/api-reference/python/mooncake-store.md +++ b/docs/source/api-reference/python/mooncake-store.md @@ -2910,6 +2910,197 @@ store.unregister_buffer(target_tensor.data_ptr()) ``` +--- + +### Session-based ranged multi-buffer transfer + +For layerwise KV load/save, resolve Master metadata once per object, then transfer +object-byte ranges across multiple buffers without re-querying Master on every layer. + +Typical flow: + +- Get: `batch_get_session_start` → `batch_get_into_multi_buffer_ranges` (per layer) → `batch_get_session_end` +- Put: `batch_put_session_start` → `batch_put_from_multi_buffer_ranges` (per layer) → `batch_put_session_end` / `batch_put_session_revoke` + +Get sessions cache a filtered `QueryResult` (single complete memory replica + lease). +Range calls only check the cached lease locally (zero Master RPCs). Put sessions +reserve object space via Master `BatchPutStart` and finalize with `BatchPutEnd`. + +Put sessions write MEMORY replicas only. `nof_replica_num > 0` is accepted only for +flexible dual-replica configs (`replica_num == 1` and `nof_replica_num == 1`), where +`batch_put_session_end` finalizes MEMORY and revokes the unused NoF reservation. +Reliable multi-replica NoF configs are rejected at session start. `end` / `revoke` +seal the session (no further range writes) and wait for in-flight range transfers +before talking to Master. + +⚠️ **Store-managed Buffer Required**: All buffers must resolve to Store-managed +registered memory before ranged zero-copy operations. + +#### batch_get_session_start() + +Query replicas once and open a get session for the given keys. + +```python +def batch_get_session_start(self, keys: List[str]) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers + +**Returns:** +- `List[int]`: Per-key status (0 = success, negative = error) + +#### batch_get_into_multi_buffer_ranges() + +Ranged get into multiple buffers using an active get session (no Master RPC). + +```python +def batch_get_into_multi_buffer_ranges( + self, + keys: List[str], + all_buffer_ptrs: List[List[int]], + all_sizes: List[List[int]], + all_src_offsets: List[List[int]], +) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers (must have an active get session) +- `all_buffer_ptrs` (List[List[int]]): Per-key list of destination buffer addresses +- `all_sizes` (List[List[int]]): Per-key list of transfer sizes in bytes +- `all_src_offsets` (List[List[int]]): Per-key list of object-byte source offsets + +**Returns:** +- `List[int]`: Bytes transferred per key (positive = success, negative = error) + +#### batch_get_session_end() + +Drop cached get-session metadata for the given keys. + +```python +def batch_get_session_end(self, keys: List[str]) -> int +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers + +**Returns:** +- `int`: 0 on success, negative on error + +#### batch_put_session_start() + +Reserve objects and open a put session without transferring data. + +```python +def batch_put_session_start( + self, + keys: List[str], + sizes: List[int], + config: ReplicateConfig = None, +) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers +- `sizes` (List[int]): Full object sizes in bytes +- `config` (ReplicateConfig, optional): Replication configuration (applies at start only). + If `group_ids` is set, its length must equal `len(keys)`. When some keys already + have a put session, they are skipped and `group_ids` is filtered to match the + remaining keys. + +**Returns:** +- `List[int]`: Per-key status (0 = success, negative = error) + +#### batch_put_from_multi_buffer_ranges() + +Ranged put from multiple buffers using an active put session (no Master RPC). + +```python +def batch_put_from_multi_buffer_ranges( + self, + keys: List[str], + all_buffer_ptrs: List[List[int]], + all_sizes: List[List[int]], + all_dst_offsets: List[List[int]], +) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers (must have an active put session) +- `all_buffer_ptrs` (List[List[int]]): Per-key list of source buffer addresses +- `all_sizes` (List[List[int]]): Per-key list of transfer sizes in bytes +- `all_dst_offsets` (List[List[int]]): Per-key list of object-byte destination offsets + +**Returns:** +- `List[int]`: Bytes transferred per key (positive = success, negative = error) + +#### batch_put_session_end() + +Finalize a put session and make objects readable. + +```python +def batch_put_session_end(self, keys: List[str]) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers + +**Returns:** +- `List[int]`: Per-key status (0 = success, negative = error) + +#### batch_put_session_revoke() + +Abort an incomplete put session and release reserved space. + +```python +def batch_put_session_revoke(self, keys: List[str]) -> List[int] +``` + +**Parameters:** +- `keys` (List[str]): Object identifiers + +**Returns:** +- `List[int]`: Per-key status (0 = success, negative = error) + +**Example:** + +
+Click to expand: Session ranged put/get example + +```python +import numpy as np + +page = 1024 +layers = 4 +keys = ["block0", "block1"] +object_sizes = [page * layers] * len(keys) + +# Prepare one registered buffer per layer for each key +src = [np.full(page, i, dtype=np.uint8) for i in range(layers)] +dst = [np.zeros(page, dtype=np.uint8) for _ in range(layers)] +for buf in src + dst: + store.register_buffer(buf.ctypes.data, buf.nbytes) + +assert all(rc == 0 for rc in store.batch_put_session_start(keys, object_sizes)) +for layer in range(layers): + ptrs = [[src[layer].ctypes.data] for _ in keys] + sizes = [[page] for _ in keys] + offsets = [[layer * page] for _ in keys] + rcs = store.batch_put_from_multi_buffer_ranges(keys, ptrs, sizes, offsets) + assert all(rc == page for rc in rcs) +assert all(rc == 0 for rc in store.batch_put_session_end(keys)) + +assert all(rc == 0 for rc in store.batch_get_session_start(keys)) +for layer in range(layers): + ptrs = [[dst[layer].ctypes.data] for _ in keys] + sizes = [[page] for _ in keys] + offsets = [[layer * page] for _ in keys] + rcs = store.batch_get_into_multi_buffer_ranges(keys, ptrs, sizes, offsets) + assert all(rc == page for rc in rcs) +assert store.batch_get_session_end(keys) == 0 +``` +
+ ## MooncakeHostMemAllocator Class The `MooncakeHostMemAllocator` class provides host memory allocation capabilities for Mooncake Store operations. diff --git a/mooncake-integration/store/store_py.cpp b/mooncake-integration/store/store_py.cpp index 166d940c5f..1d239d70b1 100644 --- a/mooncake-integration/store/store_py.cpp +++ b/mooncake-integration/store/store_py.cpp @@ -3019,6 +3019,111 @@ PYBIND11_MODULE(store, m) { "Get object data directly into multiple pre-allocated buffers for " "multiple " "keys") + .def( + "batch_get_session_start", + [](MooncakeStorePyWrapper &self, + const std::vector &keys) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_get_session_start(keys); + }, + py::arg("keys"), + "Start a get session: query replicas once and cache them") + .def( + "batch_get_into_multi_buffer_ranges", + [](MooncakeStorePyWrapper &self, + const std::vector &keys, + const std::vector> &all_buffer_ptrs, + const std::vector> &all_sizes, + const std::vector> &all_src_offsets) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_get_into_multi_buffer_ranges( + keys, CastAddrs2Ptrs(all_buffer_ptrs), all_sizes, + all_src_offsets); + }, + py::arg("keys"), py::arg("all_buffer_ptrs"), py::arg("all_sizes"), + py::arg("all_src_offsets"), + "Ranged get into multiple buffers using a get session") + .def( + "batch_get_session_end", + [](MooncakeStorePyWrapper &self, + const std::vector &keys) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return -1; + } + py::gil_scoped_release release; + return self.store_->batch_get_session_end(keys); + }, + py::arg("keys"), + "End a get session and drop cached replica metadata") + .def( + "batch_put_session_start", + [](MooncakeStorePyWrapper &self, + const std::vector &keys, + const std::vector &sizes, + const ReplicateConfig &config = ReplicateConfig{}) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_put_session_start(keys, sizes, + config); + }, + py::arg("keys"), py::arg("sizes"), + py::arg("config") = ReplicateConfig{}, + "Start a put session: reserve objects without transfer") + .def( + "batch_put_from_multi_buffer_ranges", + [](MooncakeStorePyWrapper &self, + const std::vector &keys, + const std::vector> &all_buffer_ptrs, + const std::vector> &all_sizes, + const std::vector> &all_dst_offsets) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_put_from_multi_buffer_ranges( + keys, CastAddrs2Ptrs(all_buffer_ptrs), all_sizes, + all_dst_offsets); + }, + py::arg("keys"), py::arg("all_buffer_ptrs"), py::arg("all_sizes"), + py::arg("all_dst_offsets"), + "Ranged put from multiple buffers using a put session") + .def( + "batch_put_session_end", + [](MooncakeStorePyWrapper &self, + const std::vector &keys) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_put_session_end(keys); + }, + py::arg("keys"), "Complete a put session and make objects readable") + .def( + "batch_put_session_revoke", + [](MooncakeStorePyWrapper &self, + const std::vector &keys) { + if (!self.is_client_initialized()) { + LOG(ERROR) << "Client is not initialized"; + return std::vector{}; + } + py::gil_scoped_release release; + return self.store_->batch_put_session_revoke(keys); + }, + py::arg("keys"), "Revoke an incomplete put session") .def( "get_replica_desc", [](MooncakeStorePyWrapper &self, const std::string &key) { diff --git a/mooncake-store/include/client_service.h b/mooncake-store/include/client_service.h index 7bd2e481de..78e324e8be 100644 --- a/mooncake-store/include/client_service.h +++ b/mooncake-store/include/client_service.h @@ -31,6 +31,7 @@ namespace mooncake { class PutOperation; +class RealClient; /** * @brief Result of a query operation containing replica information and lease @@ -226,6 +227,37 @@ class Client { std::vector>& batched_slices, const ReplicateConfig& config); + /** + * @brief Write slices into a memory replica at an object-byte offset. + */ + ErrorCode TransferWriteRange(const Replica::Descriptor& replica_descriptor, + std::vector& slices, + uint64_t dst_offset); + + /** + * @brief Batch ranged read against cached replicas. Fragments from all + * entries are issued as one scatter transfer so the transport can coalesce + * everything bound for the same segment, then awaited together. Requires + * memory replicas. Returns per-entry total bytes transferred or an + * ErrorCode. Used by RealClient get sessions. No Master RPC. + */ + std::vector> BatchTransferReadRanges( + const std::vector& replicas, + const std::vector>& slices, + const std::vector>& src_offsets); + + /** + * @brief Batch ranged write into cached replicas (replication). Fragments + * from all entries and all memory replicas are issued as one scatter + * transfer, then awaited together. Returns per-entry logical bytes + * transferred (counted once, not per replica) or an ErrorCode. Used by + * RealClient put sessions. + */ + std::vector> BatchTransferWriteRanges( + const std::vector>& replicas_per_entry, + const std::vector>& slices, + const std::vector>& dst_offsets); + /** * @brief Upserts data: inserts if key doesn't exist, updates if it does * @param key Object key @@ -667,6 +699,23 @@ class Client { bool IsReplicaOnLocalMemory(const Replica::Descriptor& replica); + // First half of BatchPut only (size-only slices → StartBatchPut). + // Used by RealClient put sessions; not a Master API facade. + std::vector, ErrorCode>> + StartBatchPutForSizes(const std::vector& keys, + const std::vector& object_sizes, + const ReplicateConfig& config); + + // Finalize/revoke a put session for a batch of keys. Thin wrappers over + // the master client, exposed so RealClient put sessions can end/revoke + // without touching Client internals. Same RPC path as FinalizeBatchPut. + std::vector> BatchPutEnd( + const std::vector& object_metas, + ReplicaType replica_type = ReplicaType::ALL); + std::vector> BatchPutRevoke( + const std::vector& keys, + ReplicaType replica_type = ReplicaType::ALL); + protected: /** * @brief Constructor exposed to subclasses for testing only; production @@ -704,6 +753,15 @@ class Client { ErrorCode TransferReadInternal( const Replica::Descriptor& replica_descriptor, std::vector& slices, uint64_t src_offset); + // Internal async range submission helpers (return the transfer future). + // Used by the synchronous TransferReadRange/WriteRange and the batch + // range transfer methods; not part of the public API. + std::optional SubmitRangeRead( + const Replica::Descriptor& replica_descriptor, + std::vector& slices, uint64_t src_offset); + std::optional SubmitRangeWrite( + const Replica::Descriptor& replica_descriptor, + std::vector& slices, uint64_t dst_offset); ErrorCode TransferWrite(const Replica::Descriptor& replica_descriptor, std::vector& slices); ErrorCode TransferRead(const Replica::Descriptor& replica_descriptor, diff --git a/mooncake-store/include/pyclient.h b/mooncake-store/include/pyclient.h index 4219ae6202..0ffdc8d382 100644 --- a/mooncake-store/include/pyclient.h +++ b/mooncake-store/include/pyclient.h @@ -294,6 +294,57 @@ class PyClient { const std::vector> &all_sizes, const ReplicateConfig &config = ReplicateConfig{}) = 0; + // Put/get sessions. Default stubs keep DummyClient unchanged; RealClient + // overrides with the real implementations. + virtual std::vector batch_get_session_start( + const std::vector &keys) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + virtual std::vector batch_get_into_multi_buffer_ranges( + const std::vector &keys, + const std::vector> & /*all_buffers*/, + const std::vector> & /*all_sizes*/, + const std::vector> & /*all_src_offsets*/) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + virtual int batch_get_session_end( + const std::vector & /*keys*/) { + return static_cast(toInt(ErrorCode::INVALID_PARAMS)); + } + + virtual std::vector batch_put_session_start( + const std::vector &keys, + const std::vector & /*sizes*/, + const ReplicateConfig & /*config*/ = ReplicateConfig{}) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + virtual std::vector batch_put_from_multi_buffer_ranges( + const std::vector &keys, + const std::vector> & /*all_buffers*/, + const std::vector> & /*all_sizes*/, + const std::vector> & /*all_dst_offsets*/) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + virtual std::vector batch_put_session_end( + const std::vector &keys) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + virtual std::vector batch_put_session_revoke( + const std::vector &keys) { + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + virtual std::shared_ptr get_buffer( const std::string &key) = 0; diff --git a/mooncake-store/include/real_client.h b/mooncake-store/include/real_client.h index b0f9f20006..c6f3adc69d 100644 --- a/mooncake-store/include/real_client.h +++ b/mooncake-store/include/real_client.h @@ -2,6 +2,8 @@ #include #include +#include +#include #include #include #include @@ -244,6 +246,33 @@ class RealClient : public PyClient { const std::vector> &all_sizes, const ReplicateConfig &config = ReplicateConfig{}); + std::vector batch_get_session_start( + const std::vector &keys) override; + + std::vector batch_get_into_multi_buffer_ranges( + const std::vector &keys, + const std::vector> &all_buffers, + const std::vector> &all_sizes, + const std::vector> &all_src_offsets) override; + + int batch_get_session_end(const std::vector &keys) override; + + std::vector batch_put_session_start( + const std::vector &keys, const std::vector &sizes, + const ReplicateConfig &config = ReplicateConfig{}) override; + + std::vector batch_put_from_multi_buffer_ranges( + const std::vector &keys, + const std::vector> &all_buffers, + const std::vector> &all_sizes, + const std::vector> &all_dst_offsets) override; + + std::vector batch_put_session_end( + const std::vector &keys) override; + + std::vector batch_put_session_revoke( + const std::vector &keys) override; + int put_parts(const std::string &key, std::vector> values, const ReplicateConfig &config = ReplicateConfig{}); @@ -898,6 +927,23 @@ class RealClient : public PyClient { std::unordered_map registered_buffer_sizes_; std::optional local_buffer_region_; + // KV transfer sessions (process-local; not shared with DummyClient). + // get_sessions_ stores a FilterQueryResult'd QueryResult (single complete + // memory replica + lease); ranges only compare lease locally (no Master). + // Put sessions track writable + inflight so end/revoke can seal the + // session and wait for outstanding range writes before finalize/free. + struct PutSessionEntry { + std::vector replicas; + uint64_t object_size{0}; + ReplicaWriteMode write_mode{ReplicaWriteMode::SINGLE_REPLICA}; + bool writable{true}; + size_t inflight_transfers{0}; + }; + mutable std::mutex session_mutex_; + std::condition_variable session_cv_; + std::unordered_map get_sessions_; + std::unordered_map put_sessions_; + // Dummy VA -> real VA using mapped_shms; last_hit_shm caches locality. bool map_dummy_range_in_shm(const MappedShm &shm, uint64_t dummy_addr, size_t offset, size_t size, diff --git a/mooncake-store/include/transfer_task.h b/mooncake-store/include/transfer_task.h index 8676cd64b8..bbe37968c6 100644 --- a/mooncake-store/include/transfer_task.h +++ b/mooncake-store/include/transfer_task.h @@ -565,6 +565,10 @@ class TransferSubmitter { const Replica::Descriptor& replica, std::vector& slices, uint64_t src_offset); + std::optional submitRangeWrite( + const Replica::Descriptor& replica, std::vector& slices, + uint64_t dst_offset); + TransferEngine::ScatterTransferOperation submitScatter( const std::vector& transfers); @@ -655,6 +659,10 @@ class TransferSubmitter { const AllocatedBuffer::Descriptor& handle, const std::vector& slices, uint64_t src_offset); + std::optional submitMemoryWriteOperation( + const AllocatedBuffer::Descriptor& handle, + const std::vector& slices, uint64_t dst_offset); + std::optional submitFileReadOperation( const Replica::Descriptor& replica, std::vector& slices, TransferRequest::OpCode op_code); diff --git a/mooncake-store/src/client_service.cpp b/mooncake-store/src/client_service.cpp index 2ec06bf52b..9e513d95a4 100644 --- a/mooncake-store/src/client_service.cpp +++ b/mooncake-store/src/client_service.cpp @@ -149,6 +149,71 @@ std::optional GetContiguousSliceRange( .size = total_size}; } +ErrorCode ScatterFragmentError(const Status& status) { + return status.IsInvalidArgument() ? ErrorCode::INVALID_PARAMS + : ErrorCode::TRANSFER_FAIL; +} + +// Collects the fragments of many ranged entries into a single scatter submit. +// +// One submit per entry would hand the transport one tiny transfer per key per +// layer; batching them lets the transport coalesce every fragment bound for the +// same segment into one transfer, which is what keeps layer-wise sessions at +// full bandwidth. +// +// ScatterTransferRange holds non-owning spans, so the offset/length storage +// lives here and must outlive the submit. Both vectors are reserved to the +// exact fragment count up front so pushes never reallocate under a live span. +class ScatterRangeBuilder { + public: + explicit ScatterRangeBuilder(size_t fragment_count) + : zero_offsets_(fragment_count, 0) { + remote_offsets_.reserve(fragment_count); + lengths_.reserve(fragment_count); + ranges_.reserve(fragment_count); + } + + // `error_slot` is shared by every fragment of one entry and keeps the first + // failure. Callbacks run on the waiting thread, so no locking is needed. + void Add(TransferRequest::OpCode opcode, + const AllocatedBuffer::Descriptor& handle, const Slice& slice, + uint64_t remote_offset, std::optional* error_slot) { + const size_t index = remote_offsets_.size(); + remote_offsets_.push_back(static_cast(remote_offset)); + lengths_.push_back(slice.size); + ranges_.push_back(TransferEngine::ScatterTransferRange{ + .opcode = opcode, + .remote_segment = handle.transport_endpoint_, + .remote_base_offset = handle.buffer_address_, + .remote_size = static_cast(handle.size_), + .local_buffer = slice.ptr, + .local_capacity = slice.size, + .local_offsets = std::span(&zero_offsets_[index], 1), + .remote_offsets = + std::span(&remote_offsets_[index], 1), + .lengths = std::span(&lengths_[index], 1), + .on_fragment_complete = + [error_slot](size_t, const Status& status) { + if (!status.ok() && !error_slot->has_value()) { + *error_slot = ScatterFragmentError(status); + } + }, + }); + } + + bool empty() const { return ranges_.empty(); } + + const std::vector& ranges() const { + return ranges_; + } + + private: + std::vector zero_offsets_; + std::vector remote_offsets_; + std::vector lengths_; + std::vector ranges_; +}; + struct ReplicaTransferSummary { size_t allocated_memory_replicas = 0; size_t allocated_nof_replicas = 0; @@ -2854,6 +2919,50 @@ std::vector> Client::BatchPut( return CollectResults(ops); } +std::vector, ErrorCode>> +Client::StartBatchPutForSizes(const std::vector& keys, + const std::vector& object_sizes, + const ReplicateConfig& config) { + std::vector, ErrorCode>> + results(keys.size(), tl::unexpected(ErrorCode::INVALID_PARAMS)); + if (keys.size() != object_sizes.size()) { + LOG(ERROR) << "StartBatchPutForSizes size mismatch: keys=" + << keys.size() << ", sizes=" << object_sizes.size(); + return results; + } + + ReplicateConfig client_cfg = AttachHostId(config); + if (protocol_ == "cxl") { + client_cfg.preferred_segment = local_hostname_; + } + + std::vector> batched_slices(keys.size()); + for (size_t i = 0; i < keys.size(); ++i) { + batched_slices[i] = {Slice{nullptr, object_sizes[i]}}; + } + std::vector ops = CreatePutOperations(keys, batched_slices); + StartBatchPut(ops, client_cfg); + + for (size_t i = 0; i < ops.size(); ++i) { + if (ops[i].IsResolved()) { + results[i] = tl::unexpected(ops[i].result.error()); + continue; + } + results[i] = std::move(ops[i].replicas); + } + return results; +} + +std::vector> Client::BatchPutEnd( + const std::vector& object_metas, ReplicaType replica_type) { + return master_client_.BatchPutEnd(object_metas, replica_type); +} + +std::vector> Client::BatchPutRevoke( + const std::vector& keys, ReplicaType replica_type) { + return master_client_.BatchPutRevoke(keys, replica_type); +} + tl::expected Client::Remove(const ObjectKey& key, bool force) { if (hot_cache_) { hot_cache_->BumpKeyGeneration(key); @@ -3766,16 +3875,193 @@ ErrorCode Client::TransferData(const Replica::Descriptor& replica_descriptor, return future->get(); } -ErrorCode Client::TransferReadInternal( +std::optional Client::SubmitRangeRead( const Replica::Descriptor& replica_descriptor, std::vector& slices, uint64_t src_offset) { if (!transfer_submitter_) { LOG(ERROR) << "TransferSubmitter not initialized"; - return ErrorCode::INVALID_PARAMS; + return std::nullopt; + } + return transfer_submitter_->submitRangeRead(replica_descriptor, slices, + src_offset); +} + +std::optional Client::SubmitRangeWrite( + const Replica::Descriptor& replica_descriptor, std::vector& slices, + uint64_t dst_offset) { + if (!transfer_submitter_) { + LOG(ERROR) << "TransferSubmitter not initialized"; + return std::nullopt; + } + return transfer_submitter_->submitRangeWrite(replica_descriptor, slices, + dst_offset); +} + +std::vector> Client::BatchTransferReadRanges( + const std::vector& replicas, + const std::vector>& slices, + const std::vector>& src_offsets) { + std::vector> results( + replicas.size(), tl::unexpected(ErrorCode::INVALID_PARAMS)); + if (replicas.size() != slices.size() || + replicas.size() != src_offsets.size()) { + LOG(ERROR) << "BatchTransferReadRanges size mismatch: replicas=" + << replicas.size() << ", slices=" << slices.size() + << ", offsets=" << src_offsets.size(); + return results; + } + + size_t fragment_count = 0; + for (const auto& entry : slices) { + fragment_count += entry.size(); + } + + // Every fragment of every entry goes into one scatter submit so the + // transport sees the whole layer at once instead of one transfer per key. + ScatterRangeBuilder builder(fragment_count); + std::vector> entry_errors(replicas.size()); + for (size_t i = 0; i < replicas.size(); ++i) { + if (slices[i].size() != src_offsets[i].size()) { + LOG(ERROR) << "BatchTransferReadRanges fragment count mismatch, " + << "entry=" << i << ", slices=" << slices[i].size() + << ", offsets=" << src_offsets[i].size(); + continue; // results[i] stays INVALID_PARAMS + } + if (!replicas[i].is_memory_replica()) { + LOG(ERROR) << "Range read requires a memory replica, entry=" << i; + continue; + } + + const auto& handle = + replicas[i].get_memory_descriptor().buffer_descriptor; + int64_t transferred = 0; + for (size_t j = 0; j < slices[i].size(); ++j) { + builder.Add(TransferRequest::READ, handle, slices[i][j], + src_offsets[i][j], &entry_errors[i]); + transferred += static_cast(slices[i][j].size); + } + results[i] = transferred; // optimistic; corrected on await + } + + if (builder.empty()) { + return results; + } + + auto operation = SubmitScatter(builder.ranges()); + if (!operation) { + LOG(ERROR) << "Failed to submit batch range read"; + for (auto& result : results) { + if (result.has_value()) { + result = tl::unexpected(ErrorCode::TRANSFER_FAIL); + } + } + return results; + } + (void)operation->wait(); + + for (size_t i = 0; i < results.size(); ++i) { + if (!results[i].has_value() || !entry_errors[i].has_value()) { + continue; + } + LOG(ERROR) << "Range read failed, entry=" << i + << ", error=" << static_cast(entry_errors[i].value()); + results[i] = tl::unexpected(entry_errors[i].value()); + } + return results; +} + +std::vector> Client::BatchTransferWriteRanges( + const std::vector>& replicas_per_entry, + const std::vector>& slices, + const std::vector>& dst_offsets) { + std::vector> results( + replicas_per_entry.size(), tl::unexpected(ErrorCode::INVALID_PARAMS)); + if (replicas_per_entry.size() != slices.size() || + replicas_per_entry.size() != dst_offsets.size()) { + LOG(ERROR) << "BatchTransferWriteRanges size mismatch: entries=" + << replicas_per_entry.size() << ", slices=" << slices.size() + << ", offsets=" << dst_offsets.size(); + return results; + } + + size_t fragment_count = 0; + for (size_t i = 0; i < replicas_per_entry.size(); ++i) { + for (const auto& replica : replicas_per_entry[i]) { + if (replica.is_memory_replica()) { + fragment_count += slices[i].size(); + } + } + } + + // One scatter submit covers every replica of every entry, so replication + // fans out inside a single transfer instead of one submit per fragment. + ScatterRangeBuilder builder(fragment_count); + std::vector> entry_errors( + replicas_per_entry.size()); + for (size_t i = 0; i < replicas_per_entry.size(); ++i) { + if (slices[i].size() != dst_offsets[i].size()) { + LOG(ERROR) << "BatchTransferWriteRanges fragment count mismatch, " + << "entry=" << i << ", slices=" << slices[i].size() + << ", offsets=" << dst_offsets[i].size(); + continue; // results[i] stays INVALID_PARAMS + } + + // Logical bytes: counted once per fragment, not per replica. + int64_t transferred = 0; + for (const auto& slice : slices[i]) { + transferred += static_cast(slice.size); + } + bool submitted = false; + for (const auto& replica : replicas_per_entry[i]) { + if (!replica.is_memory_replica()) { + continue; + } + const auto& handle = + replica.get_memory_descriptor().buffer_descriptor; + for (size_t j = 0; j < slices[i].size(); ++j) { + builder.Add(TransferRequest::WRITE, handle, slices[i][j], + dst_offsets[i][j], &entry_errors[i]); + submitted = true; + } + } + if (!submitted) { + results[i] = tl::unexpected(ErrorCode::INVALID_REPLICA); + continue; + } + results[i] = transferred; // optimistic; corrected on await + } + + if (builder.empty()) { + return results; + } + + auto operation = SubmitScatter(builder.ranges()); + if (!operation) { + LOG(ERROR) << "Failed to submit batch range write"; + for (auto& result : results) { + if (result.has_value()) { + result = tl::unexpected(ErrorCode::TRANSFER_FAIL); + } + } + return results; + } + (void)operation->wait(); + + for (size_t i = 0; i < results.size(); ++i) { + if (!results[i].has_value() || !entry_errors[i].has_value()) { + continue; + } + LOG(ERROR) << "Range write failed, entry=" << i + << ", error=" << static_cast(entry_errors[i].value()); + results[i] = tl::unexpected(entry_errors[i].value()); } + return results; +} - auto future = transfer_submitter_->submitRangeRead(replica_descriptor, - slices, src_offset); +ErrorCode Client::TransferReadInternal( + const Replica::Descriptor& replica_descriptor, std::vector& slices, + uint64_t src_offset) { + auto future = SubmitRangeRead(replica_descriptor, slices, src_offset); if (!future) { LOG(ERROR) << "Failed to submit range read operation"; return ErrorCode::TRANSFER_FAIL; @@ -3791,6 +4077,19 @@ ErrorCode Client::TransferWrite(const Replica::Descriptor& replica_descriptor, return TransferData(replica_descriptor, slices, TransferRequest::WRITE); } +ErrorCode Client::TransferWriteRange( + const Replica::Descriptor& replica_descriptor, std::vector& slices, + uint64_t dst_offset) { + auto future = SubmitRangeWrite(replica_descriptor, slices, dst_offset); + if (!future) { + LOG(ERROR) << "Failed to submit range write operation"; + return ErrorCode::TRANSFER_FAIL; + } + + VLOG(1) << "Using transfer strategy: " << future->strategy(); + return future->get(); +} + ErrorCode Client::TransferRead(const Replica::Descriptor& replica_descriptor, std::vector& slices) { size_t total_size = 0; diff --git a/mooncake-store/src/real_client.cpp b/mooncake-store/src/real_client.cpp index 4789f789ca..a774762a64 100644 --- a/mooncake-store/src/real_client.cpp +++ b/mooncake-store/src/real_client.cpp @@ -261,6 +261,41 @@ inline QueryResult FilterQueryResult(const QueryResult &qr, {replica}, qr.lease_timeout, include_object_checksum ? qr.object_checksum : std::nullopt); } + +// Shared object-byte range overflow check (same semantics as +// execute_ranged_read / session ranged get-put). +inline bool is_object_range_overflow(size_t offset, size_t size, size_t limit) { + return size > limit || offset > limit - size; +} + +inline const Replica::Descriptor *SelectCompleteMemoryReplica( + const std::vector &replicas, + const std::unordered_set &local_endpoints) { + const Replica::Descriptor *first_memory = nullptr; + for (const auto &r : replicas) { + if (r.status != ReplicaStatus::COMPLETE || !r.is_memory_replica()) { + continue; + } + if (local_endpoints.count(r.get_memory_descriptor() + .buffer_descriptor.transport_endpoint_)) { + return &r; + } + if (!first_memory) { + first_memory = &r; + } + } + return first_memory; +} + +inline bool HasMemoryReplica(const std::vector &replicas) { + for (const auto &r : replicas) { + if (r.is_memory_replica()) { + return true; + } + } + return false; +} + } // namespace PyClient::~PyClient() {} @@ -3290,7 +3325,8 @@ tl::expected RealClient::execute_ranged_read( return tl::unexpected(ErrorCode::INVALID_PARAMS); } size = total_size; - } else if (size > total_size || src_offset > total_size - size) { + } else if (is_object_range_overflow(src_offset, size, + static_cast(total_size))) { LOG(ERROR) << "Range overflow: src_offset=" << src_offset << " + size=" << size << " > total=" << total_size; return tl::unexpected(ErrorCode::INVALID_PARAMS); @@ -5170,7 +5206,589 @@ std::vector RealClient::batch_put_from_multi_buffers( for (const auto &result : internal_results) { results.push_back(to_py_ret(result)); } + return results; +} + +std::vector RealClient::batch_get_session_start( + const std::vector &keys) { + std::vector results( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + if (!client_) { + LOG(ERROR) << "Client is not initialized"; + return results; + } + if (keys.empty()) { + return {}; + } + + // Master interaction only here: query replicas + lease. + const auto query_results = client_->BatchQuery(keys); + auto local_endpoints = client_->GetLocalEndpoints(); + + std::lock_guard lock(session_mutex_); + for (size_t i = 0; i < keys.size(); ++i) { + if (!query_results[i]) { + results[i] = static_cast(toInt(query_results[i].error())); + get_sessions_.erase(keys[i]); + continue; + } + + auto query_result = query_results[i].value(); + if (query_result.IsLeaseExpired()) { + results[i] = static_cast(toInt(ErrorCode::LEASE_EXPIRED)); + get_sessions_.erase(keys[i]); + continue; + } + + const auto *replica = + SelectCompleteMemoryReplica(query_result.replicas, local_endpoints); + if (!replica) { + LOG(ERROR) << "No complete memory replica for key: " << keys[i]; + results[i] = static_cast(toInt(ErrorCode::INVALID_REPLICA)); + get_sessions_.erase(keys[i]); + continue; + } + + // QueryResult members are const: erase + emplace (no operator=). + get_sessions_.erase(keys[i]); + get_sessions_.emplace(keys[i], + FilterQueryResult(query_result, *replica)); + results[i] = 0; + } + return results; +} + +std::vector RealClient::batch_get_into_multi_buffer_ranges( + const std::vector &keys, + const std::vector> &all_buffers, + const std::vector> &all_sizes, + const std::vector> &all_src_offsets) { + std::vector results( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + if (!client_ || keys.size() != all_buffers.size() || + keys.size() != all_sizes.size() || + keys.size() != all_src_offsets.size()) { + LOG(ERROR) << "Invalid get ranges args"; + return results; + } + + // No Master RPC here: use cached QueryResult from session start. + // RealClient owns session state (lease/overflow checks, replica lookup); + // the actual parallel transfer is delegated to Client. + std::vector replicas; + std::vector> slices; + std::vector> src_offsets; + std::vector idx_map; // batch entry -> original key index + std::vector lease_deadlines; + + { + std::lock_guard lock(session_mutex_); + auto now = std::chrono::steady_clock::now(); + for (size_t i = 0; i < keys.size(); ++i) { + const auto &buffers = all_buffers[i]; + const auto &sizes = all_sizes[i]; + const auto &offsets = all_src_offsets[i]; + if (buffers.size() != sizes.size() || + buffers.size() != offsets.size()) { + continue; + } + auto it = get_sessions_.find(keys[i]); + if (it == get_sessions_.end()) { + continue; + } + if (it->second.IsLeaseExpired(now)) { + get_sessions_.erase(it); + results[i] = static_cast(toInt(ErrorCode::LEASE_EXPIRED)); + continue; + } + // start cached a single complete memory replica via + // FilterQueryResult. + const auto &replica = it->second.replicas.front(); + const size_t replica_limit = + replica.is_memory_replica() + ? replica.get_memory_descriptor().buffer_descriptor.size_ + : 0; + bool overflow = false; + std::vector entry_slices; + std::vector entry_offsets; + entry_slices.reserve(buffers.size()); + entry_offsets.reserve(buffers.size()); + for (size_t j = 0; j < buffers.size(); ++j) { + if (replica_limit == 0 || + is_object_range_overflow(offsets[j], sizes[j], + replica_limit)) { + overflow = true; + break; + } + entry_slices.emplace_back(Slice{buffers[j], sizes[j]}); + entry_offsets.push_back(static_cast(offsets[j])); + } + if (overflow) { + results[i] = static_cast(toInt(ErrorCode::INVALID_PARAMS)); + continue; + } + replicas.push_back(replica); + slices.push_back(std::move(entry_slices)); + src_offsets.push_back(std::move(entry_offsets)); + idx_map.push_back(i); + lease_deadlines.push_back(it->second.lease_timeout); + } + } + + if (replicas.empty()) { + return results; + } + + auto transfer = + client_->BatchTransferReadRanges(replicas, slices, src_offsets); + + // Merge results; drop sessions whose lease expired during the wait. + { + std::lock_guard lock(session_mutex_); + const auto now = std::chrono::steady_clock::now(); + for (size_t k = 0; k < transfer.size(); ++k) { + const size_t i = idx_map[k]; + if (transfer[k]) { + if (now >= lease_deadlines[k]) { + results[i] = + static_cast(toInt(ErrorCode::LEASE_EXPIRED)); + get_sessions_.erase(keys[i]); + } else { + results[i] = static_cast(transfer[k].value()); + } + } else { + results[i] = static_cast(toInt(transfer[k].error())); + } + } + } + return results; +} + +int RealClient::batch_get_session_end(const std::vector &keys) { + std::lock_guard lock(session_mutex_); + for (const auto &key : keys) { + get_sessions_.erase(key); + } + return 0; +} + +std::vector RealClient::batch_put_session_start( + const std::vector &keys, const std::vector &sizes, + const ReplicateConfig &config) { + std::vector results(keys.size(), 0); + if (!client_ || keys.size() != sizes.size()) { + LOG(ERROR) << "Invalid batch_put_session_start args"; + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + if (config.group_ids.has_value() && + config.group_ids->size() != keys.size()) { + LOG(ERROR) << "batch_put_session_start: group_ids.size()=" + << config.group_ids->size() + << ", keys.size()=" << keys.size() + << ", error=invalid_group_ids"; + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + // Session put only writes MEMORY replicas. Reliable configs that require + // NoF completion cannot be finalized correctly on this path. + const auto write_mode = DetermineReplicaWriteMode(config); + if (config.nof_replica_num > 0 && + write_mode != ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA) { + LOG(ERROR) << "batch_put_session_start rejects reliable NoF configs: " + << "session path never writes NoF " + << "(nof_replica_num=" << config.nof_replica_num + << ", replica_num=" << config.replica_num << ")"; + return std::vector( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + } + + std::vector start_keys; + std::vector start_sizes; + std::vector start_indices; + { + std::lock_guard lock(session_mutex_); + for (size_t i = 0; i < keys.size(); ++i) { + if (put_sessions_.count(keys[i]) != 0) { + LOG(ERROR) << "Put session already exists for key: " << keys[i]; + results[i] = static_cast(toInt(ErrorCode::INVALID_PARAMS)); + continue; + } + start_keys.push_back(keys[i]); + start_sizes.push_back(static_cast(sizes[i])); + start_indices.push_back(i); + } + } + if (start_keys.empty()) { + return results; + } + + // start_keys may be a subset when some keys already have sessions; keep + // group_ids aligned with the filtered key list. + ReplicateConfig start_config = config; + if (start_config.group_ids.has_value()) { + std::vector filtered_group_ids; + filtered_group_ids.reserve(start_indices.size()); + for (size_t idx : start_indices) { + filtered_group_ids.push_back(start_config.group_ids->at(idx)); + } + start_config.group_ids = std::move(filtered_group_ids); + } + + // Same first half as Client::BatchPut (StartBatchPut → master + // BatchPutStart). + auto start_responses = + client_->StartBatchPutForSizes(start_keys, start_sizes, start_config); + std::vector keys_to_revoke; + std::vector hot_keys_to_remove; + { + std::lock_guard lock(session_mutex_); + for (size_t j = 0; j < start_keys.size(); ++j) { + const size_t i = start_indices[j]; + if (!start_responses[j]) { + results[i] = + static_cast(toInt(start_responses[j].error())); + continue; + } + auto replicas = std::move(start_responses[j].value()); + if (!HasMemoryReplica(replicas)) { + hot_keys_to_remove.push_back(keys[i]); + keys_to_revoke.push_back(keys[i]); + results[i] = + static_cast(toInt(ErrorCode::INVALID_REPLICA)); + continue; + } + put_sessions_[keys[i]] = PutSessionEntry{ + .replicas = std::move(replicas), + .object_size = static_cast(sizes[i]), + .write_mode = write_mode, + }; + } + } + // Master RPCs and hot-cache updates run outside session_mutex_. + if (client_->GetHotCache() && !hot_keys_to_remove.empty()) { + client_->GetHotCache()->RemoveHotKeys(hot_keys_to_remove); + } + if (!keys_to_revoke.empty()) { + (void)client_->BatchPutRevoke(keys_to_revoke); + } + return results; +} + +std::vector RealClient::batch_put_from_multi_buffer_ranges( + const std::vector &keys, + const std::vector> &all_buffers, + const std::vector> &all_sizes, + const std::vector> &all_dst_offsets) { + std::vector results( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + if (!client_ || keys.size() != all_buffers.size() || + keys.size() != all_sizes.size() || + keys.size() != all_dst_offsets.size()) { + LOG(ERROR) << "Invalid put ranges args"; + return results; + } + + // RealClient owns session state (overflow check vs object_size, replica + // lookup); the parallel replicated transfer is delegated to Client. + std::vector> replicas_per_entry; + std::vector> slices; + std::vector> dst_offsets; + std::vector idx_map; // batch entry -> original key index + std::vector inflight_keys; + + { + std::lock_guard lock(session_mutex_); + for (size_t i = 0; i < keys.size(); ++i) { + const auto &buffers = all_buffers[i]; + const auto &sizes = all_sizes[i]; + const auto &offsets = all_dst_offsets[i]; + if (buffers.size() != sizes.size() || + buffers.size() != offsets.size()) { + continue; + } + auto it = put_sessions_.find(keys[i]); + if (it == put_sessions_.end() || !it->second.writable) { + continue; + } + bool overflow = false; + std::vector entry_slices; + std::vector entry_offsets; + entry_slices.reserve(buffers.size()); + entry_offsets.reserve(buffers.size()); + for (size_t j = 0; j < buffers.size(); ++j) { + if (is_object_range_overflow( + offsets[j], sizes[j], + static_cast(it->second.object_size))) { + overflow = true; + break; + } + entry_slices.emplace_back(Slice{buffers[j], sizes[j]}); + entry_offsets.push_back(static_cast(offsets[j])); + } + if (overflow) { + continue; // results[i] stays INVALID_PARAMS + } + ++it->second.inflight_transfers; + inflight_keys.push_back(keys[i]); + replicas_per_entry.push_back(it->second.replicas); + slices.push_back(std::move(entry_slices)); + dst_offsets.push_back(std::move(entry_offsets)); + idx_map.push_back(i); + } + } + + if (replicas_per_entry.empty()) { + return results; + } + + // Ensure end/revoke can make progress even if transfer throws. + struct InflightGuard { + RealClient *self; + std::vector keys; + ~InflightGuard() { + if (!self) { + return; + } + std::lock_guard lock(self->session_mutex_); + for (const auto &key : keys) { + auto it = self->put_sessions_.find(key); + if (it == self->put_sessions_.end()) { + continue; + } + if (it->second.inflight_transfers > 0) { + --it->second.inflight_transfers; + } + } + self->session_cv_.notify_all(); + } + } inflight_guard{this, std::move(inflight_keys)}; + + auto transfer = client_->BatchTransferWriteRanges(replicas_per_entry, + slices, dst_offsets); + + for (size_t k = 0; k < transfer.size(); ++k) { + const size_t i = idx_map[k]; + if (transfer[k]) { + results[i] = static_cast(transfer[k].value()); + } else { + results[i] = static_cast(toInt(transfer[k].error())); + } + } + return results; +} +std::vector RealClient::batch_put_session_end( + const std::vector &keys) { + std::vector results( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + if (!client_) { + LOG(ERROR) << "Client is not initialized"; + return results; + } + + // Session put only writes memory replicas (BatchTransferWriteRanges skips + // NoF). For FLEXIBLE_DUAL_REPLICA, finalize MEMORY then revoke leftover NoF + // — same split FinalizeBatchPut uses when only memory transfers succeed. + // Reliable configs with NoF are rejected at session start. + std::vector end_keys; + std::vector end_indices; + std::vector has_nof; // parallel to end_keys + { + std::unique_lock lock(session_mutex_); + for (size_t i = 0; i < keys.size(); ++i) { + auto it = put_sessions_.find(keys[i]); + if (it == put_sessions_.end()) { + results[i] = static_cast(toInt(ErrorCode::INVALID_PARAMS)); + continue; + } + bool nof = false; + for (const auto &replica : it->second.replicas) { + if (replica.is_nof_replica()) { + nof = true; + break; + } + } + // Seal first so concurrent range writes cannot start. + it->second.writable = false; + // MEMORY+NoF revoke fallback is only valid for flexible dual mode. + if (nof && it->second.write_mode != + ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA) { + LOG(ERROR) + << "batch_put_session_end: key=" << keys[i] + << " has NoF replicas under non-flexible write mode; " + << "use batch_put_session_revoke"; + results[i] = static_cast(toInt(ErrorCode::INVALID_PARAMS)); + continue; + } + end_keys.push_back(keys[i]); + end_indices.push_back(i); + has_nof.push_back(static_cast(nof)); + } + if (!end_keys.empty()) { + session_cv_.wait(lock, [&] { + for (const auto &key : end_keys) { + auto it = put_sessions_.find(key); + if (it != put_sessions_.end() && + it->second.inflight_transfers > 0) { + return false; + } + } + return true; + }); + } + } + if (end_keys.empty()) { + return results; + } + + // Session put path does not compute checksums; leave object_checksum unset. + std::vector end_metas; + end_metas.reserve(end_keys.size()); + for (const auto &key : end_keys) { + end_metas.emplace_back(ObjectMeta{key, std::nullopt}); + } + auto end_responses = client_->BatchPutEnd(end_metas, ReplicaType::MEMORY); + if (end_responses.size() != end_keys.size()) { + // Ambiguous Master state: keep sessions so caller can retry end/revoke. + for (size_t j = 0; j < end_keys.size(); ++j) { + results[end_indices[j]] = + static_cast(toInt(ErrorCode::RPC_FAIL)); + } + return results; + } + + std::vector nof_revoke_keys; + std::vector nof_revoke_end_indices; + nof_revoke_keys.reserve(end_keys.size()); + nof_revoke_end_indices.reserve(end_keys.size()); + for (size_t j = 0; j < end_keys.size(); ++j) { + const size_t i = end_indices[j]; + if (!end_responses[j]) { + // Keep session for retry/revoke; do not erase on per-key failure. + results[i] = static_cast(toInt(end_responses[j].error())); + continue; + } + if (has_nof[j]) { + nof_revoke_keys.push_back(end_keys[j]); + nof_revoke_end_indices.push_back(j); + continue; + } + { + std::lock_guard lock(session_mutex_); + put_sessions_.erase(end_keys[j]); + } + results[i] = 0; + } + + if (nof_revoke_keys.empty()) { + return results; + } + + auto nof_revoke_responses = + client_->BatchPutRevoke(nof_revoke_keys, ReplicaType::NOF_SSD); + if (nof_revoke_responses.size() != nof_revoke_keys.size()) { + for (size_t k = 0; k < nof_revoke_keys.size(); ++k) { + const size_t j = nof_revoke_end_indices[k]; + results[end_indices[j]] = + static_cast(toInt(ErrorCode::RPC_FAIL)); + // Keep session: MEMORY end may have succeeded; retry can re-end + // (idempotent) and re-revoke NoF. + } + return results; + } + for (size_t k = 0; k < nof_revoke_keys.size(); ++k) { + const size_t j = nof_revoke_end_indices[k]; + const size_t i = end_indices[j]; + if (!nof_revoke_responses[k] && + nof_revoke_responses[k].error() != ErrorCode::OBJECT_NOT_FOUND) { + results[i] = + static_cast(toInt(nof_revoke_responses[k].error())); + continue; + } + { + std::lock_guard lock(session_mutex_); + put_sessions_.erase(nof_revoke_keys[k]); + } + results[i] = 0; + } + return results; +} + +std::vector RealClient::batch_put_session_revoke( + const std::vector &keys) { + std::vector results( + keys.size(), static_cast(toInt(ErrorCode::INVALID_PARAMS))); + if (!client_) { + LOG(ERROR) << "Client is not initialized"; + return results; + } + + std::vector revoke_keys; + std::vector revoke_indices; + { + std::unique_lock lock(session_mutex_); + for (size_t i = 0; i < keys.size(); ++i) { + auto it = put_sessions_.find(keys[i]); + if (it == put_sessions_.end()) { + results[i] = static_cast(toInt(ErrorCode::INVALID_PARAMS)); + continue; + } + // Seal session and wait for outstanding range writes before free. + it->second.writable = false; + revoke_keys.push_back(keys[i]); + revoke_indices.push_back(i); + } + if (!revoke_keys.empty()) { + session_cv_.wait(lock, [&] { + for (const auto &key : revoke_keys) { + auto it = put_sessions_.find(key); + if (it != put_sessions_.end() && + it->second.inflight_transfers > 0) { + return false; + } + } + return true; + }); + } + } + if (revoke_keys.empty()) { + return results; + } + + // Same path FinalizeBatchPut uses on transfer failure. + if (client_->GetHotCache() && !revoke_keys.empty()) { + client_->GetHotCache()->RemoveHotKeys(revoke_keys); + } + auto revoke_responses = + client_->BatchPutRevoke(revoke_keys, ReplicaType::ALL); + std::lock_guard lock(session_mutex_); + if (revoke_responses.size() != revoke_keys.size()) { + // Ambiguous Master state: keep sessions so caller can retry revoke. + for (size_t j = 0; j < revoke_keys.size(); ++j) { + results[revoke_indices[j]] = + static_cast(toInt(ErrorCode::RPC_FAIL)); + } + return results; + } + for (size_t j = 0; j < revoke_keys.size(); ++j) { + const size_t i = revoke_indices[j]; + if (!revoke_responses[j]) { + if (revoke_responses[j].error() == ErrorCode::OBJECT_NOT_FOUND) { + // Nothing left on Master; drop the local session. + put_sessions_.erase(revoke_keys[j]); + results[i] = 0; + } else { + // Keep session for retry. + results[i] = + static_cast(toInt(revoke_responses[j].error())); + } + continue; + } + put_sessions_.erase(revoke_keys[j]); + results[i] = 0; + } return results; } diff --git a/mooncake-store/src/transfer_task.cpp b/mooncake-store/src/transfer_task.cpp index 134472b256..e4b55f1822 100644 --- a/mooncake-store/src/transfer_task.cpp +++ b/mooncake-store/src/transfer_task.cpp @@ -1290,6 +1290,25 @@ std::optional TransferSubmitter::submitMemoryReadOperation( return std::nullopt; } +std::optional TransferSubmitter::submitMemoryWriteOperation( + const AllocatedBuffer::Descriptor& handle, const std::vector& slices, + uint64_t dst_offset) { + TransferStrategy strategy = selectStrategy(handle, slices); + + if (strategy == TransferStrategy::LOCAL_MEMCPY) { + return submitMemcpyOperation(handle, slices, TransferRequest::WRITE, + dst_offset); + } + if (strategy == TransferStrategy::TRANSFER_ENGINE) { + return submitTransferEngineOperation( + handle, slices, TransferRequest::WRITE, dst_offset); + } + + LOG(ERROR) << "Write only supports LOCAL_MEMCPY or TRANSFER_ENGINE, got: " + << strategy; + return std::nullopt; +} + std::optional TransferSubmitter::submitRangeRead( const Replica::Descriptor& replica, std::vector& slices, uint64_t src_offset) { @@ -1301,7 +1320,8 @@ std::optional TransferSubmitter::submitRangeRead( size_t slices_size = 0; for (const auto& s : slices) slices_size += s.size; - if (src_offset + slices_size > handle.size_) { + if (src_offset > std::numeric_limits::max() - slices_size || + src_offset + slices_size > handle.size_) { LOG(ERROR) << "Range read overflow: src_offset=" << src_offset << " + slices_size=" << slices_size << " > handle.size_=" << handle.size_; @@ -1325,6 +1345,42 @@ std::optional TransferSubmitter::submitRangeRead( return future; } +std::optional TransferSubmitter::submitRangeWrite( + const Replica::Descriptor& replica, std::vector& slices, + uint64_t dst_offset) { + std::optional future; + + if (replica.is_memory_replica()) { + auto& mem_desc = replica.get_memory_descriptor(); + auto& handle = mem_desc.buffer_descriptor; + + size_t slices_size = 0; + for (const auto& s : slices) slices_size += s.size; + if (dst_offset > std::numeric_limits::max() - slices_size || + dst_offset + slices_size > handle.size_) { + LOG(ERROR) << "Range write overflow: dst_offset=" << dst_offset + << " + slices_size=" << slices_size + << " > handle.size_=" << handle.size_; + return std::nullopt; + } + + future = submitMemoryWriteOperation(handle, slices, dst_offset); + } else if (replica.is_nof_replica()) { + LOG(ERROR) << "Range write not supported for NoF replicas"; + return std::nullopt; + } else if (replica.is_disk_replica() || replica.is_local_disk_replica()) { + LOG(ERROR) + << "Range write not supported for disk replicas (use full write)"; + return std::nullopt; + } + + if (future.has_value()) { + updateTransferMetrics(slices, TransferRequest::WRITE); + } + + return future; +} + #ifdef USE_NOF std::optional TransferSubmitter::submitSpdkNofOperation( const AllocatedBuffer::Descriptor& handle, void* ptr, size_t size, diff --git a/mooncake-store/tests/client_buffer_test.cpp b/mooncake-store/tests/client_buffer_test.cpp index d144f3c561..56b8445bb9 100644 --- a/mooncake-store/tests/client_buffer_test.cpp +++ b/mooncake-store/tests/client_buffer_test.cpp @@ -67,9 +67,9 @@ TEST_F(ClientBufferTest, SpdkDmaAllocatorDestroysWithSpdkFree) { std::shared_ptr allocator; try { - allocator = ClientBufferAllocator::create( - buffer_size, "tcp", /*use_hugepage=*/false, - /*use_spdk_dma=*/true); + allocator = ClientBufferAllocator::create(buffer_size, "tcp", + /*use_hugepage=*/false, + /*use_spdk_dma=*/true); } catch (const std::bad_alloc&) { GTEST_SKIP() << "SPDK DMA allocation is unavailable in this environment"; diff --git a/mooncake-store/tests/e2e/run_session_ranges_tcp_e2e.sh b/mooncake-store/tests/e2e/run_session_ranges_tcp_e2e.sh new file mode 100755 index 0000000000..92648ee830 --- /dev/null +++ b/mooncake-store/tests/e2e/run_session_ranges_tcp_e2e.sh @@ -0,0 +1,79 @@ +#!/usr/bin/env bash +# Real TCP transport e2e for put/get session ranged multi-buffer APIs. +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "$SCRIPT_DIR/../../.." && pwd) +BUILD_DIR=${BUILD_DIR:-"$REPO_ROOT/build"} +LOG_DIR=${LOG_DIR:-/tmp/mooncake_session_ranges_tcp_e2e} +MASTER_RPC=${MASTER_RPC:-127.0.0.1:50051} +MASTER_HOST=${MASTER_RPC%:*} +MASTER_PORT=${MASTER_RPC##*:} +LEASE_TTL_MS=${LEASE_TTL_MS:-60000} + +MASTER_BIN="$BUILD_DIR/mooncake-store/src/mooncake_master" +STORE_SO=$(ls "$BUILD_DIR"/mooncake-integration/store.cpython-*.so 2>/dev/null | head -1 || true) + +MASTER_PID="" + +cleanup() { + if [[ -n "${MASTER_PID}" ]]; then + kill "${MASTER_PID}" >/dev/null 2>&1 || true + wait "${MASTER_PID}" 2>/dev/null || true + fi +} +trap cleanup EXIT + +rm -rf "$LOG_DIR" +mkdir -p "$LOG_DIR" + +if [[ ! -x "$MASTER_BIN" ]]; then + echo "missing mooncake_master at $MASTER_BIN; build target mooncake_master first" >&2 + exit 2 +fi +if [[ -z "$STORE_SO" ]]; then + echo "missing store python module under $BUILD_DIR/mooncake-integration" >&2 + exit 3 +fi + +pkill -f "$MASTER_BIN" >/dev/null 2>&1 || true +sleep 1 + +echo "starting mooncake_master protocol=tcp rpc=$MASTER_RPC" +"$MASTER_BIN" \ + --rpc_address="$MASTER_HOST" \ + --rpc_port="$MASTER_PORT" \ + --default_kv_lease_ttl="$LEASE_TTL_MS" \ + --enable_http_metadata_server=false \ + --logtostderr=true \ + >"$LOG_DIR/master.log" 2>&1 & +MASTER_PID=$! +sleep 2 + +if ! kill -0 "$MASTER_PID" >/dev/null 2>&1; then + echo "mooncake_master failed to start; see $LOG_DIR/master.log" >&2 + exit 4 +fi + +export PYTHONPATH="$BUILD_DIR/mooncake-integration${PYTHONPATH:+:$PYTHONPATH}" +export LD_LIBRARY_PATH="$BUILD_DIR/mooncake-store/src:$BUILD_DIR/mooncake-common:$BUILD_DIR/mooncake-transfer-engine/src${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +export MOONCAKE_PROTOCOL=tcp +export MOONCAKE_DEVICE= +export MOONCAKE_MASTER="$MASTER_RPC" +export MOONCAKE_TE_META_DATA_SERVER=P2PHANDSHAKE +export MOONCAKE_LOCAL_HOSTNAME=localhost:17814 +# Ascend-enabled builds still try to install ascend transport unless forced. +export MC_FORCE_TCP="${MC_FORCE_TCP:-1}" + +set +e +python3 "$SCRIPT_DIR/session_ranges_tcp_e2e.py" | tee "$LOG_DIR/client.log" +RC=${PIPESTATUS[0]} +set -e + +if [[ "$RC" -eq 0 ]]; then + echo "PASSED: session ranges TCP e2e (logs: $LOG_DIR)" +else + echo "FAILED: session ranges TCP e2e rc=$RC (logs: $LOG_DIR)" >&2 + tail -n 80 "$LOG_DIR/master.log" >&2 || true +fi +exit "$RC" diff --git a/mooncake-store/tests/e2e/session_ranges_tcp_e2e.py b/mooncake-store/tests/e2e/session_ranges_tcp_e2e.py new file mode 100644 index 0000000000..de15499851 --- /dev/null +++ b/mooncake-store/tests/e2e/session_ranges_tcp_e2e.py @@ -0,0 +1,152 @@ +#!/usr/bin/env python3 +"""TCP e2e for put/get session ranged multi-buffer APIs. + +Requires a running mooncake_master and built store Python module. + +Example: + PYTHONPATH=build/mooncake-integration \\ + MOONCAKE_PROTOCOL=tcp \\ + MOONCAKE_MASTER=127.0.0.1:50051 \\ + MOONCAKE_TE_META_DATA_SERVER=P2PHANDSHAKE \\ + python3 mooncake-store/tests/e2e/session_ranges_tcp_e2e.py +""" + +from __future__ import annotations + +import ctypes +import os +import sys +import time + + +def _require_store(): + try: + import store # type: ignore + except Exception as exc: # pragma: no cover + print(f"import_fail {exc}", flush=True) + raise SystemExit(10) + return store + + +def _ptr(buf: ctypes.Array) -> int: + return ctypes.addressof(buf) + + +def run() -> int: + store = _require_store() + + master = os.getenv("MOONCAKE_MASTER", "127.0.0.1:50051") + metadata = os.getenv("MOONCAKE_TE_META_DATA_SERVER", "P2PHANDSHAKE") + protocol = os.getenv("MOONCAKE_PROTOCOL", "tcp") + device = os.getenv("MOONCAKE_DEVICE", "") + hostname = os.getenv("MOONCAKE_LOCAL_HOSTNAME", "localhost:17814") + segment = int(os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", str(64 * 1024 * 1024))) + local_buf = int(os.getenv("MOONCAKE_LOCAL_BUFFER_SIZE", str(64 * 1024 * 1024))) + + num_layers = int(os.getenv("E2E_NUM_LAYERS", "4")) + page_size = int(os.getenv("E2E_PAGE_SIZE", "4096")) + num_keys = int(os.getenv("E2E_NUM_KEYS", "3")) + object_size = page_size * num_layers + + print( + f"e2e_session_ranges protocol={protocol} master={master} " + f"keys={num_keys} layers={num_layers} page={page_size}", + flush=True, + ) + + mc = store.MooncakeDistributedStore() + setup_ret = mc.setup( + hostname, metadata, segment, local_buf, protocol, device, master + ) + print(f"setup_ret {setup_ret}", flush=True) + if setup_ret != 0: + return setup_ret + + src = (ctypes.c_char * (object_size * num_keys))() + dst = (ctypes.c_char * (object_size * num_keys))() + for i in range(len(src)): + src[i] = ord("a") + (i % 26) + dst[i] = ord("B") + + assert mc.register_buffer(_ptr(src), len(src)) == 0 + assert mc.register_buffer(_ptr(dst), len(dst)) == 0 + + keys = [f"session_e2e_key_{i}_{int(time.time())}" for i in range(num_keys)] + sizes = [object_size] * num_keys + + put_start = mc.batch_put_session_start(keys, sizes) + print(f"batch_put_session_start {put_start}", flush=True) + if any(rc != 0 for rc in put_start): + return 20 + + for layer in range(num_layers): + all_buffers = [] + all_sizes = [] + all_dst_offsets = [] + for i in range(num_keys): + offset = i * object_size + layer * page_size + all_buffers.append([_ptr(src) + offset]) + all_sizes.append([page_size]) + all_dst_offsets.append([layer * page_size]) + put_rcs = mc.batch_put_from_multi_buffer_ranges( + keys, all_buffers, all_sizes, all_dst_offsets + ) + print(f"batch_put_ranges layer={layer} rcs={put_rcs}", flush=True) + if any(rc != page_size for rc in put_rcs): + mc.batch_put_session_revoke(keys) + return 21 + + put_end = mc.batch_put_session_end(keys) + print(f"batch_put_session_end {put_end}", flush=True) + if any(rc != 0 for rc in put_end): + return 22 + + get_start = mc.batch_get_session_start(keys) + print(f"batch_get_session_start {get_start}", flush=True) + if any(rc != 0 for rc in get_start): + return 30 + + for layer in range(num_layers): + all_buffers = [] + all_sizes = [] + all_src_offsets = [] + for i in range(num_keys): + offset = i * object_size + layer * page_size + all_buffers.append([_ptr(dst) + offset]) + all_sizes.append([page_size]) + all_src_offsets.append([layer * page_size]) + get_rcs = mc.batch_get_into_multi_buffer_ranges( + keys, all_buffers, all_sizes, all_src_offsets + ) + print(f"batch_get_ranges layer={layer} rcs={get_rcs}", flush=True) + if any(rc != page_size for rc in get_rcs): + mc.batch_get_session_end(keys) + return 31 + + get_end = mc.batch_get_session_end(keys) + print(f"batch_get_session_end {get_end}", flush=True) + if get_end != 0: + return 32 + + if bytes(src) != bytes(dst): + print("data_mismatch", flush=True) + return 40 + + # Revoke path: start put then revoke before end. + revoke_keys = [f"session_e2e_revoke_{int(time.time())}"] + revoke_start = mc.batch_put_session_start(revoke_keys, [page_size]) + print(f"batch_put_session_start_revoke {revoke_start}", flush=True) + if any(rc != 0 for rc in revoke_start): + return 50 + revoke_rcs = mc.batch_put_session_revoke(revoke_keys) + print(f"batch_put_session_revoke {revoke_rcs}", flush=True) + if any(rc != 0 for rc in revoke_rcs): + return 51 + + print("session_ranges_tcp_e2e PASSED", flush=True) + mc.close() + return 0 + + +if __name__ == "__main__": + sys.exit(run()) diff --git a/mooncake-store/tests/pybind_client_test.cpp b/mooncake-store/tests/pybind_client_test.cpp index b9c7057097..2183649284 100644 --- a/mooncake-store/tests/pybind_client_test.cpp +++ b/mooncake-store/tests/pybind_client_test.cpp @@ -941,6 +941,486 @@ TEST_F(RealClientTest, TestBatchPutAndGetMultiBuffers) { << "Dst data buffer unregistration should succeed"; } +TEST_F(RealClientTest, TestPutGetSessionRanges) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17814", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kNumLayers = 4; + constexpr size_t kPageSize = 64; + constexpr size_t kObjectSize = kPageSize * kNumLayers; + constexpr size_t kNumKeys = 3; + + std::string src_data(kObjectSize * kNumKeys, 'A'); + std::string dst_data(kObjectSize * kNumKeys, 'B'); + for (size_t i = 0; i < src_data.size(); ++i) { + src_data[i] = static_cast('a' + (i % 26)); + } + + ASSERT_EQ(py_client_->register_buffer(src_data.data(), src_data.size()), 0); + ASSERT_EQ(py_client_->register_buffer(dst_data.data(), dst_data.size()), 0); + + std::vector keys; + std::vector object_sizes; + keys.reserve(kNumKeys); + object_sizes.reserve(kNumKeys); + for (size_t i = 0; i < kNumKeys; ++i) { + keys.push_back("session_key_" + std::to_string(i)); + object_sizes.push_back(kObjectSize); + } + + auto put_start_rcs = + py_client_->batch_put_session_start(keys, object_sizes); + ASSERT_EQ(put_start_rcs.size(), kNumKeys); + for (auto rc : put_start_rcs) { + EXPECT_EQ(rc, 0) << "batch_put_session_start should succeed"; + } + + for (size_t layer = 0; layer < kNumLayers; ++layer) { + std::vector> all_buffers(kNumKeys); + std::vector> all_sizes(kNumKeys); + std::vector> all_dst_offsets(kNumKeys); + for (size_t i = 0; i < kNumKeys; ++i) { + char* layer_ptr = + src_data.data() + i * kObjectSize + layer * kPageSize; + all_buffers[i] = {layer_ptr}; + all_sizes[i] = {kPageSize}; + all_dst_offsets[i] = {layer * kPageSize}; + } + auto put_rcs = py_client_->batch_put_from_multi_buffer_ranges( + keys, all_buffers, all_sizes, all_dst_offsets); + ASSERT_EQ(put_rcs.size(), kNumKeys); + for (auto rc : put_rcs) { + EXPECT_EQ(rc, static_cast(kPageSize)); + } + } + + auto put_end_rcs = py_client_->batch_put_session_end(keys); + ASSERT_EQ(put_end_rcs.size(), kNumKeys); + for (auto rc : put_end_rcs) { + EXPECT_EQ(rc, 0) << "batch_put_session_end should succeed"; + } + + auto get_start_rcs = py_client_->batch_get_session_start(keys); + ASSERT_EQ(get_start_rcs.size(), kNumKeys); + for (auto rc : get_start_rcs) { + EXPECT_EQ(rc, 0) << "batch_get_session_start should succeed"; + } + + for (size_t layer = 0; layer < kNumLayers; ++layer) { + std::vector> all_buffers(kNumKeys); + std::vector> all_sizes(kNumKeys); + std::vector> all_src_offsets(kNumKeys); + for (size_t i = 0; i < kNumKeys; ++i) { + char* layer_ptr = + dst_data.data() + i * kObjectSize + layer * kPageSize; + all_buffers[i] = {layer_ptr}; + all_sizes[i] = {kPageSize}; + all_src_offsets[i] = {layer * kPageSize}; + } + auto get_rcs = py_client_->batch_get_into_multi_buffer_ranges( + keys, all_buffers, all_sizes, all_src_offsets); + ASSERT_EQ(get_rcs.size(), kNumKeys); + for (auto rc : get_rcs) { + EXPECT_EQ(rc, static_cast(kPageSize)); + } + } + + EXPECT_EQ(py_client_->batch_get_session_end(keys), 0); + EXPECT_EQ(dst_data, src_data); + + // Revoke path: start then revoke without end. + std::vector revoke_keys = {"session_revoke_0"}; + std::vector revoke_sizes = {kObjectSize}; + auto revoke_start = + py_client_->batch_put_session_start(revoke_keys, revoke_sizes); + ASSERT_EQ(revoke_start.size(), 1); + EXPECT_EQ(revoke_start[0], 0); + auto revoke_rcs = py_client_->batch_put_session_revoke(revoke_keys); + ASSERT_EQ(revoke_rcs.size(), 1); + EXPECT_EQ(revoke_rcs[0], 0); + auto missing = py_client_->batch_get_session_start(revoke_keys); + ASSERT_EQ(missing.size(), 1); + EXPECT_LT(missing[0], 0); + + ASSERT_EQ(py_client_->unregister_buffer(src_data.data()), 0); + ASSERT_EQ(py_client_->unregister_buffer(dst_data.data()), 0); +} + +// Abnormal put/get session cases. See check table in PR / review notes. +TEST_F(RealClientTest, TestPutGetSessionAbnormal) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17815", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kPage = 64; + constexpr size_t kObjectSize = kPage * 2; + const int kInvalidParams = + static_cast(toInt(ErrorCode::INVALID_PARAMS)); + + std::string buf(kObjectSize, 'x'); + ASSERT_EQ(py_client_->register_buffer(buf.data(), buf.size()), 0); + + // --- Put: ranges/end/revoke without start --- + { + std::vector keys = {"no_put_session"}; + auto ranges = py_client_->batch_put_from_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage}}, {{0}}); + ASSERT_EQ(ranges.size(), 1); + EXPECT_EQ(ranges[0], kInvalidParams); + + auto end_rcs = py_client_->batch_put_session_end(keys); + ASSERT_EQ(end_rcs.size(), 1); + EXPECT_EQ(end_rcs[0], kInvalidParams); + + auto revoke_rcs = py_client_->batch_put_session_revoke(keys); + ASSERT_EQ(revoke_rcs.size(), 1); + EXPECT_EQ(revoke_rcs[0], kInvalidParams); + } + + // --- Put: keys/sizes mismatch --- + { + auto rcs = + py_client_->batch_put_session_start({"a", "b"}, {kObjectSize}); + ASSERT_EQ(rcs.size(), 2); + EXPECT_EQ(rcs[0], kInvalidParams); + EXPECT_EQ(rcs[1], kInvalidParams); + } + + // --- Put: duplicate start --- + { + std::vector keys = {"dup_put"}; + auto first = py_client_->batch_put_session_start(keys, {kObjectSize}); + ASSERT_EQ(first.size(), 1); + EXPECT_EQ(first[0], 0); + auto second = py_client_->batch_put_session_start(keys, {kObjectSize}); + ASSERT_EQ(second.size(), 1); + EXPECT_EQ(second[0], kInvalidParams); + EXPECT_EQ(py_client_->batch_put_session_revoke(keys)[0], 0); + } + + // --- Put: range overflow past object_size --- + { + std::vector keys = {"overflow_put"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kObjectSize})[0], + 0); + auto overflow = py_client_->batch_put_from_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage}}, + {{kObjectSize}}); // offset == size + ASSERT_EQ(overflow.size(), 1); + EXPECT_EQ(overflow[0], kInvalidParams); + EXPECT_EQ(py_client_->batch_put_session_revoke(keys)[0], 0); + } + + // --- Put: buffer/size/offset arity mismatch --- + { + std::vector keys = {"arity_put"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kObjectSize})[0], + 0); + auto bad = py_client_->batch_put_from_multi_buffer_ranges( + keys, {{buf.data(), buf.data() + kPage}}, {{kPage}}, {{0}}); + ASSERT_EQ(bad.size(), 1); + EXPECT_EQ(bad[0], kInvalidParams); + EXPECT_EQ(py_client_->batch_put_session_revoke(keys)[0], 0); + } + + // --- Put: end clears session; second end fails --- + { + std::vector keys = {"put_end_once"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kObjectSize})[0], + 0); + ASSERT_EQ(py_client_->batch_put_from_multi_buffer_ranges( + keys, {{buf.data()}}, {{kObjectSize}}, {{0}})[0], + static_cast(kObjectSize)); + EXPECT_EQ(py_client_->batch_put_session_end(keys)[0], 0); + EXPECT_EQ(py_client_->batch_put_session_end(keys)[0], kInvalidParams); + // ranges after end also fail + EXPECT_EQ(py_client_->batch_put_from_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage}}, {{0}})[0], + kInvalidParams); + } + + // --- Get: ranges without start --- + { + std::vector keys = {"no_get_session"}; + auto ranges = py_client_->batch_get_into_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage}}, {{0}}); + ASSERT_EQ(ranges.size(), 1); + EXPECT_EQ(ranges[0], kInvalidParams); + } + + // --- Get: start on missing object --- + { + auto rcs = py_client_->batch_get_session_start({"missing_obj"}); + ASSERT_EQ(rcs.size(), 1); + EXPECT_LT(rcs[0], 0); + } + + // --- Get: end clears session; ranges after end fail --- + { + std::vector keys = {"put_end_once"}; // exists from above + auto start = py_client_->batch_get_session_start(keys); + ASSERT_EQ(start.size(), 1); + EXPECT_EQ(start[0], 0); + EXPECT_EQ(py_client_->batch_get_session_end(keys), 0); + auto after_end = py_client_->batch_get_into_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage}}, {{0}}); + ASSERT_EQ(after_end.size(), 1); + EXPECT_EQ(after_end[0], kInvalidParams); + } + + // --- Get: arity mismatch while session exists --- + { + std::vector keys = {"put_end_once"}; + ASSERT_EQ(py_client_->batch_get_session_start(keys)[0], 0); + auto bad = py_client_->batch_get_into_multi_buffer_ranges( + keys, {{buf.data()}}, {{kPage, kPage}}, {{0}}); + ASSERT_EQ(bad.size(), 1); + EXPECT_EQ(bad[0], kInvalidParams); + EXPECT_EQ(py_client_->batch_get_session_end(keys), 0); + } + + ASSERT_EQ(py_client_->unregister_buffer(buf.data()), 0); +} + +// Put session must survive BatchPutEnd failure so caller can retry end/revoke. +TEST_F(RealClientTest, TestPutSessionKeptAfterEndFailure) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17817", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kSize = 128; + const int kObjectNotFound = + static_cast(toInt(ErrorCode::OBJECT_NOT_FOUND)); + const int kInvalidParams = + static_cast(toInt(ErrorCode::INVALID_PARAMS)); + + std::vector keys = {"session_end_fail_keep"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kSize})[0], 0); + + // Drop Master reservation while keeping the local put session. + ASSERT_NE(py_client_->client_, nullptr); + auto revoked = py_client_->client_->BatchPutRevoke(keys, ReplicaType::ALL); + ASSERT_EQ(revoked.size(), 1u); + ASSERT_TRUE(revoked[0].has_value()); + + auto end1 = py_client_->batch_put_session_end(keys); + ASSERT_EQ(end1.size(), 1u); + EXPECT_EQ(end1[0], kObjectNotFound) + << "end should surface Master OBJECT_NOT_FOUND"; + + // Session must still exist: second end is not INVALID_PARAMS. + auto end2 = py_client_->batch_put_session_end(keys); + ASSERT_EQ(end2.size(), 1u); + EXPECT_EQ(end2[0], kObjectNotFound) + << "put session should be kept after end failure"; + + // Revoke clears the local session even when Master object is already gone. + auto revoke_rcs = py_client_->batch_put_session_revoke(keys); + ASSERT_EQ(revoke_rcs.size(), 1u); + EXPECT_EQ(revoke_rcs[0], 0); + EXPECT_EQ(py_client_->batch_put_session_end(keys)[0], kInvalidParams); + EXPECT_EQ(py_client_->batch_put_session_revoke(keys)[0], kInvalidParams); +} + +// Reliable NoF configs cannot be finalized by the session put path. +TEST_F(RealClientTest, TestPutSessionStartRejectsReliableNofConfig) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17819", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + const int kInvalidParams = + static_cast(toInt(ErrorCode::INVALID_PARAMS)); + ReplicateConfig config; + config.replica_num = 1; + config.nof_replica_num = 2; // RELIABLE_MULTI_REPLICA + + auto rcs = + py_client_->batch_put_session_start({"reliable_nof"}, {128}, config); + ASSERT_EQ(rcs.size(), 1u); + EXPECT_EQ(rcs[0], kInvalidParams); +} + +// Filtered start_keys must keep group_ids aligned (skip existing sessions). +TEST_F(RealClientTest, TestPutSessionStartFiltersGroupIds) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17820", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kSize = 128; + const int kInvalidParams = + static_cast(toInt(ErrorCode::INVALID_PARAMS)); + + ReplicateConfig first_config; + first_config.group_ids = std::vector{"group_a"}; + ASSERT_EQ(py_client_->batch_put_session_start({"group_key_a"}, {kSize}, + first_config)[0], + 0); + + ReplicateConfig batch_config; + batch_config.group_ids = std::vector{"group_a", "group_b"}; + auto rcs = py_client_->batch_put_session_start( + {"group_key_a", "group_key_b"}, {kSize, kSize}, batch_config); + ASSERT_EQ(rcs.size(), 2u); + EXPECT_EQ(rcs[0], kInvalidParams) << "existing session should be skipped"; + EXPECT_EQ(rcs[1], 0) << "filtered group_ids must match remaining keys"; + + auto revoke_rcs = + py_client_->batch_put_session_revoke({"group_key_a", "group_key_b"}); + ASSERT_EQ(revoke_rcs.size(), 2u); + EXPECT_EQ(revoke_rcs[0], 0); + EXPECT_EQ(revoke_rcs[1], 0); +} + +// Session put finalizes MEMORY only; complete object must be readable. +TEST_F(RealClientTest, TestPutSessionEndCompletesMemoryReplica) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17818", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kSize = 256; + std::string src(kSize, 'P'); + std::string dst(kSize, 'Q'); + ASSERT_EQ(py_client_->register_buffer(src.data(), src.size()), 0); + ASSERT_EQ(py_client_->register_buffer(dst.data(), dst.size()), 0); + + std::vector keys = {"session_memory_end"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kSize})[0], 0); + ASSERT_EQ(py_client_->batch_put_from_multi_buffer_ranges( + keys, {{src.data()}}, {{kSize}}, {{0}})[0], + static_cast(kSize)); + ASSERT_EQ(py_client_->batch_put_session_end(keys)[0], 0); + + auto descs = py_client_->get_replica_desc(keys[0]); + ASSERT_EQ(descs.size(), 1u); + EXPECT_TRUE(descs[0].is_memory_replica()); + EXPECT_EQ(descs[0].status, ReplicaStatus::COMPLETE); + + ASSERT_EQ(py_client_->batch_get_session_start(keys)[0], 0); + ASSERT_EQ(py_client_->batch_get_into_multi_buffer_ranges( + keys, {{dst.data()}}, {{kSize}}, {{0}})[0], + static_cast(kSize)); + EXPECT_EQ(py_client_->batch_get_session_end(keys), 0); + EXPECT_EQ(dst, src); + + ASSERT_EQ(py_client_->unregister_buffer(src.data()), 0); + ASSERT_EQ(py_client_->unregister_buffer(dst.data()), 0); +} + +// Lease expire on get ranges must drop the cached get session. +TEST_F(RealClientTest, TestGetSessionLeaseExpiredDropsSession) { + constexpr uint64_t kLeaseTtlMs = 50; + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder() + .set_default_kv_lease_ttl(kLeaseTtlMs) + .build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ASSERT_EQ( + py_client_->setup_real("localhost:17816", "P2PHANDSHAKE", + 16 * 1024 * 1024, 16 * 1024 * 1024, + FLAGS_protocol, rdma_devices, master_address_), + 0); + + constexpr size_t kSize = 128; + const int kLeaseExpired = static_cast(toInt(ErrorCode::LEASE_EXPIRED)); + const int kInvalidParams = + static_cast(toInt(ErrorCode::INVALID_PARAMS)); + + std::string src(kSize, 'S'); + std::string dst(kSize, 'D'); + ASSERT_EQ(py_client_->register_buffer(src.data(), src.size()), 0); + ASSERT_EQ(py_client_->register_buffer(dst.data(), dst.size()), 0); + + std::vector keys = {"lease_session_key"}; + ASSERT_EQ(py_client_->batch_put_session_start(keys, {kSize})[0], 0); + ASSERT_EQ(py_client_->batch_put_from_multi_buffer_ranges( + keys, {{src.data()}}, {{kSize}}, {{0}})[0], + static_cast(kSize)); + ASSERT_EQ(py_client_->batch_put_session_end(keys)[0], 0); + + ASSERT_EQ(py_client_->batch_get_session_start(keys)[0], 0); + + // Wait past cached lease_deadline (client-local check; no Master query). + std::this_thread::sleep_for(std::chrono::milliseconds(kLeaseTtlMs + 50)); + + auto expired = py_client_->batch_get_into_multi_buffer_ranges( + keys, {{dst.data()}}, {{kSize}}, {{0}}); + ASSERT_EQ(expired.size(), 1); + EXPECT_EQ(expired[0], kLeaseExpired) + << "get ranges after lease ttl should return LEASE_EXPIRED"; + + // Session must have been erased: next ranges sees no session. + auto again = py_client_->batch_get_into_multi_buffer_ranges( + keys, {{dst.data()}}, {{kSize}}, {{0}}); + ASSERT_EQ(again.size(), 1); + EXPECT_EQ(again[0], kInvalidParams) + << "get session should be dropped after LEASE_EXPIRED"; + + // get_end is still safe (idempotent erase). + EXPECT_EQ(py_client_->batch_get_session_end(keys), 0); + + ASSERT_EQ(py_client_->unregister_buffer(src.data()), 0); + ASSERT_EQ(py_client_->unregister_buffer(dst.data()), 0); +} + TEST_F(RealClientTest, TestBatchAndNormalGetReplicaDesc) { // Start in-proc master ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) From 8c6095c06e20848506cbf91ef4a714924e7b03b1 Mon Sep 17 00:00:00 2001 From: Copilot <198982749+Copilot@users.noreply.github.com> Date: Mon, 10 Aug 2026 14:20:32 +0800 Subject: [PATCH 016/483] [Store] Support custom SSH ports in SPDK target creation (#3307) * Initial plan * Add configurable SPDK SSH port Co-authored-by: stmatengss <11641725+stmatengss@users.noreply.github.com> --------- Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com> Co-authored-by: stmatengss <11641725+stmatengss@users.noreply.github.com> --- .../ssd/nvmf-ssd-deployment-guide.md | 1 + mooncake-wheel/mooncake/spdk_tgt_create.py | 33 ++++++++++++++--- mooncake-wheel/tests/test_spdk_tgt_create.py | 37 +++++++++++++++++++ 3 files changed, 65 insertions(+), 6 deletions(-) create mode 100644 mooncake-wheel/tests/test_spdk_tgt_create.py diff --git a/docs/source/deployment/ssd/nvmf-ssd-deployment-guide.md b/docs/source/deployment/ssd/nvmf-ssd-deployment-guide.md index 5b27e077b1..56745af781 100644 --- a/docs/source/deployment/ssd/nvmf-ssd-deployment-guide.md +++ b/docs/source/deployment/ssd/nvmf-ssd-deployment-guide.md @@ -126,6 +126,7 @@ python3 -m mooncake.spdk_tgt_create \ | `ip` | IP address of the target node. | | `path` | SPDK installation path on the target node. | | `pci` | PCI addresses of SSDs to register with the target. Use commas to separate multiple PCI addresses. If this field is omitted, SPDK-ready or unmounted NVMe devices on the target node are registered. | +| `--port` | SSH port used to connect to target nodes. The default value is `22`. | | `--core-mask` | CPU core mask used to start `nvmf_tgt` with `-m`. The default value is `0xff`. | **Tip**: Run `/path/scripts/setup.sh status` on a target node to list available PCI addresses. diff --git a/mooncake-wheel/mooncake/spdk_tgt_create.py b/mooncake-wheel/mooncake/spdk_tgt_create.py index f5d8554369..f8bd0e528d 100644 --- a/mooncake-wheel/mooncake/spdk_tgt_create.py +++ b/mooncake-wheel/mooncake/spdk_tgt_create.py @@ -40,9 +40,16 @@ class SPDKTgtCreator: 'buf_cache_size': '-b', } - def __init__(self, spdk_targets: List[str], transport_options: Dict[str, Any] = None, core_mask: str = '0xff'): + def __init__( + self, + spdk_targets: List[str], + transport_options: Dict[str, Any] = None, + core_mask: str = '0xff', + port: int = 22, + ): self.spdk_targets = spdk_targets self.core_mask = core_mask + self.port = port self.transport_options = dict(self.DEFAULT_TRANSPORT_OPTIONS) if transport_options: self.transport_options.update(transport_options) @@ -120,7 +127,14 @@ def _parse_spdk_targets(self) -> List[Dict[str, Any]]: return target_configs - def _ssh_connect(self, ip: str, username: str = 'root', password: str = None, key_file: str = None) -> paramiko.SSHClient: + def _ssh_connect( + self, + ip: str, + username: str = 'root', + password: str = None, + key_file: str = None, + port: int = 22, + ) -> paramiko.SSHClient: """ Establish an SSH connection to the target host. """ @@ -130,10 +144,10 @@ def _ssh_connect(self, ip: str, username: str = 'root', password: str = None, ke try: if key_file: self.logger.info(f"Connecting to {ip} using key file {key_file}") - ssh.connect(ip, username=username, key_filename=key_file) + ssh.connect(ip, port=port, username=username, key_filename=key_file) else: self.logger.info(f"Connecting to {ip} using password authentication") - ssh.connect(ip, username=username, password=password) + ssh.connect(ip, port=port, username=username, password=password) return ssh except Exception as e: self.logger.error(f"Failed to connect to {ip}: {e}") @@ -390,7 +404,7 @@ def deploy_target(self, target_config: Dict[str, Any]) -> bool: try: # Establish SSH connection - ssh = self._ssh_connect(ip) + ssh = self._ssh_connect(ip, port=self.port) try: auto_discovered = not pci_devices @@ -486,6 +500,8 @@ def parse_arguments(): help='Number of shared buffers reserved for each poll group (default: 32)') parser.add_argument('--username', type=str, default='root', help='SSH username for target nodes (default: root)') + parser.add_argument('--port', type=int, default=22, + help='SSH port for target nodes (default: 22)') parser.add_argument('--password', type=str, help='SSH password for target nodes') parser.add_argument('--key-file', type=str, @@ -508,7 +524,12 @@ def main(): } try: - creator = SPDKTgtCreator(args.spdk_target_info, transport_options, args.core_mask) + creator = SPDKTgtCreator( + args.spdk_target_info, + transport_options, + args.core_mask, + args.port, + ) success = creator.deploy_all_targets() exit(0 if success else 1) except Exception as e: diff --git a/mooncake-wheel/tests/test_spdk_tgt_create.py b/mooncake-wheel/tests/test_spdk_tgt_create.py new file mode 100644 index 0000000000..4e9c6c001e --- /dev/null +++ b/mooncake-wheel/tests/test_spdk_tgt_create.py @@ -0,0 +1,37 @@ +import sys +import unittest +from unittest.mock import patch + +try: + import paramiko +except ModuleNotFoundError: + raise unittest.SkipTest("paramiko is required for SPDK target tests") + +from mooncake.spdk_tgt_create import SPDKTgtCreator, parse_arguments + + +def test_parse_arguments_accepts_ssh_port(): + with patch.object(sys, "argv", [ + "spdk_tgt_create", + "--spdk_target_info", + "ip:127.0.0.1 path:/home/spdk", + "--port", + "2222", + ]): + args = parse_arguments() + + assert args.port == 2222 + + +def test_ssh_connect_uses_requested_port(): + creator = SPDKTgtCreator(["ip:127.0.0.1 path:/home/spdk"]) + + with patch("mooncake.spdk_tgt_create.paramiko.SSHClient") as ssh_client: + creator._ssh_connect("127.0.0.1", port=2222) + + ssh_client.return_value.connect.assert_called_once_with( + "127.0.0.1", + port=2222, + username="root", + password=None, + ) From 626de3f4e7619231be23ebf247ae1d8fd3284214 Mon Sep 17 00:00:00 2001 From: Zhanhao Cao Date: Mon, 10 Aug 2026 17:33:29 +0800 Subject: [PATCH 017/483] [Doc] Update Mooncake PG design documentation (#3353) --- .../source/api-reference/python/ep-backend.md | 120 ++-- docs/source/design/index.md | 2 +- docs/source/design/mooncake-backend-pg.md | 614 +++++++++++++----- docs/source/design/mooncake-ep.md | 4 +- .../troubleshooting/pg-ep-troubleshooting.md | 79 +-- 5 files changed, 544 insertions(+), 275 deletions(-) diff --git a/docs/source/api-reference/python/ep-backend.md b/docs/source/api-reference/python/ep-backend.md index a3c12a38dd..f4c3e04ead 100644 --- a/docs/source/api-reference/python/ep-backend.md +++ b/docs/source/api-reference/python/ep-backend.md @@ -1,14 +1,13 @@ -# Mooncake EP & Mooncake Backend (PG) +# Mooncake EP & Mooncake PG ## Overview Mooncake provides two closely related components for fault-tolerant MoE inference: -- **Mooncake Backend (PG)** is a `torch.distributed` ProcessGroup backend. It - registers the `mooncake` accelerator backend and the `mooncake-cpu` backend, - implements common collective and point-to-point APIs, tracks active ranks, and - exposes elastic recovery helpers. +- **Mooncake PG** is a `torch.distributed` ProcessGroup backend. It registers + the `mooncake` accelerator backend and the `mooncake-cpu` backend, implements + collective and point-to-point APIs, and exposes dynamic-membership helpers. - **Mooncake EP** is an expert-parallel dispatch/combine runtime for latency-sensitive MoE inference. It follows the DeepEP low-latency programming model while adding rank activeness awareness and Mooncake transport support. @@ -17,8 +16,9 @@ The usual integration pattern is to initialize a Mooncake process group first, then construct a Mooncake EP `Buffer` from that group. The process group is used both for regular collectives and for exchanging EP bootstrap metadata. -For implementation details, see the [Mooncake Backend (PG) design guide](../../design/mooncake-backend-pg.md) -and the [Mooncake EP design guide](../../design/mooncake-ep.md). +For implementation details, see the +[Mooncake PG design guide](../../design/mooncake-backend-pg.md) and the +[Mooncake EP design guide](../../design/mooncake-ep.md). ## Installation and build notes @@ -35,7 +35,7 @@ match the active `torch.__version__`. If the current PyTorch version does not match a built extension, import will fail with a message such as `Mooncake PG was not built against torch==...`. -## Mooncake Backend (PG) quick start +## Mooncake PG quick start ### CUDA backend @@ -54,14 +54,10 @@ local_rank = int(os.environ.get("LOCAL_RANK", rank)) torch.cuda.set_device(local_rank) device = torch.device("cuda", local_rank) -# Backend-level active-rank mask. Use int32 and place it on the backend device. -active_ranks = torch.ones(world_size, dtype=torch.int32, device=device) - dist.init_process_group( backend="mooncake", rank=rank, world_size=world_size, - pg_options=pg.MooncakeBackendOptions(active_ranks), ) x = torch.tensor([rank + 1], dtype=torch.int32, device=device) @@ -77,18 +73,20 @@ torchrun --nproc-per-node=2 pg_quickstart.py ### CPU backend -Use `backend="mooncake-cpu"` and put `active_ranks` on CPU: +Use `backend="mooncake-cpu"`: ```python -active_ranks = torch.ones(world_size, dtype=torch.int32) dist.init_process_group( backend="mooncake-cpu", rank=rank, world_size=world_size, - pg_options=pg.MooncakeBackendOptions(active_ranks), ) ``` +`pg_options` is optional for a fixed-size group using the default failure +handling. Pass `MooncakeBackendOptions` when reserving additional group +capacity, joining as an extension, or selecting a non-default failure mode. + ### Selecting network devices To explicitly restrict Mooncake to a list of NIC / HCA devices, call @@ -103,27 +101,38 @@ pg.set_device_filter(["mlx5_1", "mlx5_2"]) For test and benchmark commands, the same setting is commonly passed through `MOONCAKE_PGTEST_DEVICE_FILTERS=mlx5_1,mlx5_2`. -## Mooncake Backend (PG) API reference +## Mooncake PG Torch API reference ### `MooncakeBackendOptions` ```python +pg.MooncakeBackendOptions(max_group_size) +pg.MooncakeBackendOptions(max_group_size, is_extension) +pg.MooncakeBackendOptions( + max_group_size, + is_extension, + auto_deactivate_on_failure, + auto_sync_on_failure, +) + +# Explicit active-rank mirror overloads pg.MooncakeBackendOptions(active_ranks) pg.MooncakeBackendOptions(active_ranks, is_extension) -pg.MooncakeBackendOptions(active_ranks, is_extension, max_world_size) +pg.MooncakeBackendOptions(active_ranks, is_extension, max_group_size) ``` Arguments: -- `active_ranks`: `torch.int32` tensor used as the backend-level rank-health - mask. For `mooncake`, it must be on the accelerator device; for - `mooncake-cpu`, it must be on CPU. When `max_world_size` is set, size this - tensor to `max_world_size`, not the current visible world size. +- `max_group_size`: fixed in-group slot capacity. It must be at least the + initially declared group size and cannot be increased later. +- `active_ranks`: optional contiguous `torch.int32` storage used as a mirror of + committed PG membership. Its initial contents are ignored. Size it to + `max_group_size`; it may be on CPU or GPU. - `is_extension`: set to `True` for a replacement or joining process that will enter an existing group through `join_group()`. -- `max_world_size`: optional upper bound for reserved rank slots. It lets - healthy ranks reserve inactive future ranks while keeping - `dist.get_world_size()` equal to the current active size. +- `auto_deactivate_on_failure` and `auto_sync_on_failure`: select automatic or + framework-managed failure handling. Both default to `True`; auto-sync + requires auto-deactivation. ### Utility functions @@ -133,15 +142,17 @@ Arguments: | `pg.set_device_filter(filters)` | Restrict NIC/HCA selection. | Call before `init_process_group()`. | | `pg.set_transfer_engine(engine)` | Reuse an external `TransferEngine`. | The engine must outlive all process groups. | | `pg.get_active_ranks(backend)` | Return the backend active-rank tensor. | Used by EP fallback and recovery paths. | -| `pg.get_num_synced_ranks(backend)` | Return the number of ranks synchronized by the backend. | Diagnostic helper. | -| `pg.extend_group_size_to(backend, size)` | Reserve additional inactive ranks. | Newly extended ranks do not participate until recovered. | -| `pg.get_peer_state(backend, ranks)` | Check whether candidate ranks have published peer metadata. | Collective among healthy ranks. | -| `pg.recover_ranks(backend, ranks)` | Activate ready ranks and publish extension state. | Requires peer metadata to be ready. | -| `pg.join_group(backend)` | Joiner-side blocking call for extension ranks. | Used after `is_extension=True` initialization. | +| `pg.get_num_synced_ranks(backend)` | Return the number of locally activatable group slots. | Diagnostic helper. | +| `pg.get_peer_state(backend, ranks)` | Read locally mirrored activation readiness. | A lightweight, communication-free query. | +| `pg.activate_ranks(backend, ranks)` | Propose activation through the Coordinator. | A single call from any online rank is sufficient. | +| `pg.recover_ranks(backend, ranks)` | Propose activation through the Coordinator. | Compatibility alias to `activate_ranks`. | +| `pg.deactivate_ranks(backend, ranks)` | Propose deactivation through the Coordinator. | A single call from any online rank is sufficient. | +| `pg.join_group(backend)` | Confirm readiness for activation and remain blocked until activation actually occurs. | Used for scale-up, replacement, and in-place rejoin. | +| `pg.sync_after_failure(backend)` | Report current link observations, wait for reconciliation, and apply the latest group view. | Called automatically when `auto_sync_on_failure=True`; it may also be called manually. | ### Supported distributed operations -Mooncake Backend implements the following `torch.distributed` APIs. Support may +Mooncake PG implements the following `torch.distributed` APIs. Support may depend on device type, dtype, PyTorch version, and whether the current backend is `mooncake` or `mooncake-cpu`; run the PG tests on the target environment before production use. @@ -150,29 +161,27 @@ production use. | --- | --- | --- | | Collectives | `all_reduce`, `broadcast`, `all_gather`, `all_gather_into_tensor`, `reduce_scatter_tensor`, `all_to_all`, `barrier`, `reduce`, `gather`, `scatter` | Active ranks participate; inactive ranks are skipped by backend internals. | | Async work | `dist.all_reduce(..., async_op=True)` | Wait on the returned work object, then synchronize the device stream as needed. | -| P2P | `isend`, `irecv`, `batch_isend_irecv` | Single-tensor P2P is routed through the Mooncake P2P shim. | +| P2P | `isend`, `irecv`, `batch_isend_irecv` | Single-tensor P2P is routed through the Mooncake backend shim. | ## Elastic recovery protocol -Mooncake PG supports a two-sided recovery protocol. Existing healthy ranks poll -for replacement rank readiness, then activate those ranks. Replacement ranks -start in extension mode, publish metadata, and wait until healthy ranks recover -them. +Mooncake PG separates join preparation from membership activation. A joining or +recovering rank completes local warmup, calls `join_group()`, and waits. An +existing rank may poll local readiness and then issue the activation proposal. +The Coordinator validates and distributes the resulting membership. ### Healthy-rank side ```python from mooncake import pg -active_ranks = torch.tensor([1, 1, 0], dtype=torch.int32, device=device) dist.init_process_group( backend="mooncake", rank=rank, world_size=2, pg_options=pg.MooncakeBackendOptions( - active_ranks, + 3, # max_group_size False, # is_extension - 3, # max_world_size ), ) @@ -191,31 +200,35 @@ pg.recover_ranks(backend, join_ranks) ```python from mooncake import pg -active_ranks = torch.tensor([1, 1, 1], dtype=torch.int32, device=device) dist.init_process_group( backend="mooncake", rank=2, world_size=3, pg_options=pg.MooncakeBackendOptions( - active_ranks, + 3, # max_group_size True, # is_extension - 3, # max_world_size ), ) backend = dist.group.WORLD + +# Collectives are local-only before join_group. Use this +# window for framework-specific preparation, for example: +# capture_cuda_graphs() +# warm_up_model() + pg.join_group(backend) ``` Important semantics: -- `get_peer_state()` is collective among the current healthy ranks. Call it in a - consistent order across those ranks. -- New ranks are inactive after `extend_group_size_to()` and become collective - participants only after `recover_ranks()`. -- A joining rank initialized with `is_extension=True` starts with local-only - behavior and blocks in `join_group()` until the corresponding healthy ranks - publish recovery state. +- `get_peer_state()` is a local best-effort readiness query, not a collective. +- Capacity must be reserved with `max_group_size` when founding members create + the group. A joining registration appends inactive slots within that capacity. +- A joining rank starts with local-only collective behavior until `join_group`. + The join call then blocks until a Coordinator-approved activation commits. +- A single `activate_ranks()` call, or its `recover_ranks()` alias, from any + online rank is sufficient; redundant equivalent calls are safe. - Subgroups must be created in the same order on healthy and joining processes, following PyTorch `new_group()` ordering rules. @@ -403,15 +416,16 @@ updates rank activeness so EP transport metadata and QPs can be refreshed. There are two active-rank tensors in the API surface: -- **PG active-rank mask**: passed to `pg.MooncakeBackendOptions`. This is the - backend-level health mask used by collective and recovery logic. +- **PG active-rank mask**: passed to `pg.MooncakeBackendOptions`. This mirrors + the Coordinator's committed membership. - **EP active-rank tensor**: passed to `Buffer.dispatch()` and `Buffer.combine()`. It is also rank-level (`[num_ranks]`, `torch.int32`) and may be updated by EP kernels when timeout detection marks a peer as failed. -In simple integrations these tensors often carry the same health information, -but they are passed through different API layers. Keep their dtype, device, and -shape consistent with the process group world size or reserved `max_world_size`. +Their values may coincide in a simple integration, but their semantics are not +interchangeable: PG membership is configuration, while EP may update its mask +from kernel-level timeout observations. Keep the mapping, dtype, device, and +capacity consistent when propagating committed PG membership into EP. ## Tests and examples diff --git a/docs/source/design/index.md b/docs/source/design/index.md index 6b69421b4e..507f5cbd9d 100644 --- a/docs/source/design/index.md +++ b/docs/source/design/index.md @@ -30,7 +30,7 @@ distributed execution components. | Document | Description | |----------|-------------| -| [Mooncake Backend (PG)](mooncake-backend-pg) | Fault-tolerant PyTorch process-group backend. | +| [Mooncake PG](mooncake-backend-pg) | Elastic PyTorch process group. | | [Mooncake EP](mooncake-ep) | Expert-parallel communication and recovery. | | [TENT](tent/overview) | Next-generation transfer engine design. | | [TENT Benchmark](tent/tebench) | TENT benchmark framework and methodology. | diff --git a/docs/source/design/mooncake-backend-pg.md b/docs/source/design/mooncake-backend-pg.md index ed61f776a3..630b3ea5ba 100644 --- a/docs/source/design/mooncake-backend-pg.md +++ b/docs/source/design/mooncake-backend-pg.md @@ -1,210 +1,521 @@ -# Mooncake Backend (PG) Design +# Mooncake PG Design -Mooncake Backend is a `torch.distributed` ProcessGroup backend for Mooncake. It -provides collective and point-to-point communication primitives, rank-health -tracking, and elastic recovery hooks for inference systems that need to keep -serving after partial rank failures. +Mooncake PG is a communication library built on Mooncake Transfer Engine. It +provides collective and P2P operations together with dynamic membership, fault +tolerance, and recovery. -This document is intended for developers who maintain Mooncake PG itself or -integrate it into higher-level serving systems. +## Dynamic Membership at a Glance -## Goals +If you are used to a fixed-membership process group, a natural mental model is +that a group is simply _the set of ranks that were there when the group was +created_. If the set changes, you create another group. -Mooncake Backend is designed to: +Mooncake PG uses a slightly different model. Think of a group as a **dinner +table with numbered seats**. -- integrate with PyTorch through the standard ProcessGroup extension mechanism; -- expose `mooncake` for accelerator tensors and `mooncake-cpu` for CPU tensors; -- support common collective APIs used by inference engines; -- track active and inactive ranks so collectives can continue after failures; -- allow replacement ranks to publish metadata, join an existing group, and be - activated by healthy ranks; -- reuse Mooncake Transfer Engine and topology information for data movement. +The table has a fixed number of seats, decided before dinner starts. During the +meal, however, not every seat has to be occupied. -Non-goals: +A diner may leave, but the seat does not disappear. The other diners do not +shuffle around to fill the gap; nobody gets a new seat number just because +someone stepped away. Later, that diner may return to the same seat, or someone +else may occupy it. New diners may also occupy seats that were reserved but +never used. -- It is not a drop-in replacement for every NCCL/Gloo behavior. Validate each - collective, dtype, and topology required by the application. -- Elastic recovery is an explicit protocol. The backend does not silently add a - new process to all collectives without application coordination. +So there are really two separate questions: -## Relationship with `torch.distributed` +- **How many seats does the table have?** +- **Which seats are currently occupied and participating in the dinner?** -Mooncake registers two PyTorch backends when the PG extension module is imported: +The first stays fixed. The second may change over time. -- `mooncake-cpu`, registered for CPU devices; -- `mooncake`, registered for accelerator devices such as CUDA or MUSA depending - on the build. +That is the basic intuition behind Mooncake PG's dynamic membership. -The backend class itself derives from `c10d::ProcessGroup`. Applications use -regular PyTorch APIs such as `dist.init_process_group()`, `dist.all_reduce()`, -`dist.new_group()`, and `dist.batch_isend_irecv()`. +### From the dinner table to Mooncake PG -Point-to-point dispatch in PyTorch expects a `c10d::Backend` object. Mooncake PG -therefore includes a lightweight P2P shim that delegates `send` and `recv` calls -back to the owning `MooncakeBackend` instance. +In Mooncake PG, the numbered seats correspond to **rank slots**. A group +reserves a fixed number of rank slots when it is created, bounded by +`max_group_size`. Which of those ranks currently participate in group +operations is tracked separately by `active_ranks`. -## Main runtime objects +For example, a group can reserve eight rank slots while initially activating +only four -- four diners at an eight-seat table: -### `MooncakeBackendOptions` +``` +max_group_size = 8 +active_ranks = [1, 1, 1, 1, 0, 0, 0, 0] +``` + +If rank 2 is later deactivated -- one diner leaves the table: + +``` +active_ranks = [1, 1, 0, 1, 0, 0, 0, 0] +``` + +its rank slot remains reserved, and the other ranks are not renumbered -- the +empty seat simply remains empty. + +Rank 4 may also subsequently become active without filling that hole: + +``` +active_ranks = [1, 1, 0, 1, 1, 0, 0, 0] +``` + +The group keeps the same rank structure throughout these changes; only the set +of active participants changes. + +This design also extends naturally to fault tolerance and recovery. A failed +rank can be removed from the active membership without reshaping the group, +while a recovered rank can later rejoin through the same membership mechanism. + +## Core Concepts -`MooncakeBackendOptions` carries Mooncake-specific process-group configuration: +### Capacity and size -| Field | Meaning | -| --- | --- | -| `activeRanks_` | Rank-health tensor exposed to collectives and user code. | -| `isExtension_` | Whether this process is a joining/replacement rank. | -| `maxWorldSize_` | Optional reserved capacity for future ranks. | +Mooncake PG has capacity at two scopes: -The `activeRanks_` tensor must be `torch.int32`. It must be on CPU for -`mooncake-cpu` and on the accelerator device for `mooncake`. When -`maxWorldSize_` is set, `activeRanks_` must be sized to `maxWorldSize_` so the -backend can reserve inactive rank slots. +- Process-level `max_world_size` bounds the number of ranks in the whole world. +- Group-level `max_group_size` bounds the number of ranks in a group. -### Transfer group metadata +Despite the name, a group's `size` is not the number of active members. It is +one past the highest active in-group rank and therefore acts as a rank-index +upper bound. For example: -Each backend owns shared metadata for: +```text +active_ranks = [1, 0, 1, 0] +size = 3 +max_group_size = 4 +``` + +User code that allocates rank-indexed buffers must cover at least `size` +entries, including inactive holes below that bound. Per-group masks such as +`active_ranks` and `failed_ranks_hint` span `max_group_size`. + +An extension rank before `joinGroup` (`Isolated` or `Quiescing`) is the +exception: `getSize()` continues to report the group's declared size for +PyTorch rank validation, even though its effective membership is temporarily +`{self}`. + +Holes preserve rank numbers. So user code that needs the number of participants +must count `active_ranks` rather than use `getSize()` or +`dist.get_world_size()`. + +### Rank namespaces + +A process has one `GlobalRank` in the world and may have a different +`InGroupRank` in each group: + +- `GlobalRank` indexes the world-capacity namespace `[0, max_world_size)`. +- `InGroupRank` indexes the group-capacity namespace `[0, max_group_size)`. -- the current rank and backend index; -- current capacity (`size`) and visible active size (`activeSize`); -- host and device active-rank masks; -- peer connection state; -- rank-local and rank-global mapping; -- store handles and extension state used by recovery; -- P2P proxy and connection poller state. +The distinction matters whenever a group contains only part of the world or +orders its ranks differently. Group-level operations, including +`activate_ranks` and `deactivate_ranks`, use in-group ranks unless stated +otherwise. -`size` is the reserved capacity. `activeSize` is the visible group size returned -by `dist.get_world_size()`. With `max_world_size`, `size` may be larger than -`activeSize`; inactive slots are masked out by the active-rank state. +### Rank state -### Transfer Engine ownership +The Coordinator tracks a process-level `RankState`: -By default, Mooncake Backend initializes its own Transfer Engine. Advanced -integrations may call `pg.set_transfer_engine(engine)` before -`init_process_group()` to inject an external Transfer Engine. In that mode, the -caller owns the engine and must keep it alive until all Mooncake process groups -using it are destroyed. +| State | Meaning | +| --------- | ----------------------------------------------------------------------------------------------- | +| `Offline` | No usable Agent session exists for this global rank. | +| `Synced` | The Agent session is synchronized, but the rank is not in the current healthy set. | +| `Healthy` | The Agent session is synchronized and current link evidence places the rank in the healthy set. | -## Initialization lifecycle +A new Agent session moves a rank from `Offline` to `Synced`. Link evidence can +promote it to `Healthy` or demote it back to `Synced`; losing the Agent session +makes it `Offline`. -The initialization flow is: +The Coordinator combines link reports from all ranks to determine which ranks +are `Healthy`. The healthy ranks must be connected to one another in both +directions. A registered rank that does not meet these conditions remains +`Synced`. -1. Python imports `mooncake.pg`, which loads a PyTorch-version-specific native - extension. -2. The extension registers `mooncake` / `mooncake-cpu` with PyTorch. -3. `dist.init_process_group()` invokes the backend factory with PyTorch - distributed options and optional `MooncakeBackendOptions`. -4. The backend initializes active-rank masks and reserved rank slots. -5. Non-extension ranks publish local peer metadata and wait until current peers - are connected. -6. Extension ranks enter local-only mode and wait for the explicit join protocol. +`RankState` is independent of any group. It describes whether a process is +ready for data-plane communication; participation in a particular group is +described separately by group membership. -The important distinction is that reserving capacity does not automatically make -future ranks active. New ranks are masked until `recover_ranks()` activates them. +### Group membership -## Active ranks and dynamic world size +`rank_order` connects the two namespaces: it records the `GlobalRank` assigned +to each `InGroupRank`. This is a stable slot assignment, not a list of current +participants. Whether an assigned rank participates is recorded separately by +its `GroupMemberState` and reflected in `active_ranks`. -Mooncake PG tracks two related concepts: +For example, suppose global ranks 2, 5, and 7 form a group with +`max_group_size = 5`. Inside that group they are in-group ranks 0, 1, and 2, so +`rank_order = [2, 5, 7]`. Calling `deactivate_ranks([1])` targets global rank 5: +its member state becomes `Inactive`, `active_ranks` becomes +`[1, 0, 1, 0, 0]`, and the rank order does not change. -- **Reserved capacity** (`size`): how many rank slots the backend knows about. -- **Visible active size** (`activeSize`): the current group size visible through - PyTorch APIs. +If global rank 9 is later added, it is assigned in-group rank 3 and the rank +order becomes `[2, 5, 7, 9]`. Existing in-group ranks are never renumbered; +new ranks extend the order, while unused group capacity remains unassigned. -When `max_world_size` is larger than the initial `world_size`, Mooncake reserves -extra slots but marks them inactive. This lets healthy ranks poll for joiner -metadata and activate the joiners later without reconstructing the process group. +`GroupMemberState` has the following values: -`extend_group_size_to(size)` can also increase capacity. Newly extended ranks -start inactive; the application must call `get_peer_state()` and -`recover_ranks()` before they participate in collectives. +| State | Meaning | +| -------------------- | ------------------------------------------------------------------------ | +| `None` | The rank has not registered with this group. | +| `Inactive` | The rank is registered but has not declared itself ready for activation. | +| `AwaitingActivation` | The rank calls `join_group` and declares itself ready for activation. | +| `Active` | The rank participates in collective operations. | +| `Left` | The rank unregistered from this group. | -## Elastic recovery protocol +Founding members become `Active` directly during group bootstrap. A joining +member's activation follows `Inactive` → `AwaitingActivation` → `Active`. +Deactivation returns an active member to `Inactive`. -Mooncake PG uses a two-phase protocol for recovery and scale-up: +A `Healthy` rank is ready for data-plane communication, but it participates in +a group only when its group member state is `Active`. + +Membership changes are checked against both kinds of state. To activate ranks, +the group must be ready, every newly activated target must be +`AwaitingActivation`, `Healthy`, and have a published endpoint, and every rank +in the resulting active set must be mutually connected with every other rank in +that set. An early activation request remains pending until these conditions +hold or its admission timeout expires. + +## Architecture + +The data plane executes transfers directly through Mooncake Transfer Engine, +while the control plane tracks processes, connectivity, endpoints, and +committed membership. ```mermaid -sequenceDiagram - participant H as Healthy ranks - participant J as Joining rank - participant S as Store / metadata - - H->>S: init_process_group(world_size=M, max_world_size=N) - J->>S: init_process_group(world_size=N, is_extension=True) - J->>S: publish local peer metadata - H->>S: get_peer_state(join_ranks) - H->>S: recover_ranks(join_ranks) - S-->>J: extension state - J->>J: join_group() returns - H->>J: collectives include recovered ranks +flowchart LR + Framework[Framework] --> Torch[torch.distributed] + + subgraph PG[Mooncake PG] + Comm[Communicator] + Agent[Agent] + Coordinator[Coordinator on rank 0] + Workers[Collective worker and P2P proxy] + Comm <--> Agent + Agent <--> Coordinator + Comm --> Workers + end + + Torch --> Comm + Workers --> TE[Transfer Engine] +``` + +### Process context and communicators + +Each process has one Mooncake PG context. It owns or references the Transfer +Engine and hosts an Agent and the process-wide worker managers; global rank 0 +also hosts the Coordinator. + +### Coordinator and Agents + +The Coordinator owns process `RankState` and each group's member states. +It serializes membership changes and publishes them as `GroupView`s. Agents +mirror those views and apply them to local communicators. Collective and P2P +workers report link evidence; the Coordinator derives the mutually connected +healthy set and decides the resulting state changes. + +The Coordinator currently runs on global rank 0 and is not highly available. + +### Collectives + +Collectives use a direct-write design over Transfer Engine. Each communicator +registers their send, receive, and synchronization buffers. +The worker records transfer failures in `failed_ranks_hint` and reports to the +Agent; membership is not changed on the collective worker side. + +### P2P + +P2P send and receive use a receiver-driven, credit-based protocol. A receive +operation reserves chunks from a receive pool and writes `CreditSlot`s to the +sender. Each credit identifies the destination chunk and length. The sender +then stages the corresponding data in a send-pool chunk, performs a TE write to +that destination, and writes an `AckSlot` back. The receiver copies the +acknowledged chunk into the user buffer and returns the chunk to the pool. + +Each communicator has per-peer operation queues and separate credit and +acknowledgement rings. The control slots carry a group epoch and sequence +number, and a matching header/footer token prevents a partially written slot +from being consumed. Epoch checks discard control traffic left over from an +older group view. + +The send and receive polling threads and their fixed-size chunk pools are +shared by all communicators on the same device. Chunk allocation never blocks +a polling thread. When a pool has no free chunk, the operation remains pending +and the poller retries it later. +A transfer error or timeout resets the affected peer lane and reports failure +evidence through the same control-plane path as a collective failure. + +## Planned scaling + +This section describes planned scaling through Mooncake PG dynamic membership. + +### Scale-up + +The group must have enough unused `max_group_size` capacity. A joining process +declares an extended rank order (also `is_extension=True` in the PyTorch +integration). The existing group adopts the appended slots as inactive; this +step alone never activates them. + +The later join and activation steps are: + +1. The joining rank starts in `Isolated` with an effective `{self}` + membership. This gives the upper layer framework a local-only window for + initialization and warmup before the rank can affect existing members. + Collectives in this state must not be interpreted as results from the + eventual group. +2. When local preparation is complete, the joining rank calls `join_group`. + The call enters `Quiescing`, drains previously issued collective and P2P + work, marks the member `AwaitingActivation`, and waits. +3. Any online rank calls `activate_ranks`. The Coordinator admits the request + only after the activation conditions described above hold for the complete + future active set. It distributes the new membership and waits for the + required ranks to apply it; both `join_group` and `activate_ranks` then + return. + +An activation request may arrive before `join_group`; it waits for the joining +rank to become ready rather than bypassing the checks. + +### Scale-down + +After the upper layer framework stops issuing operations that use the old +membership, any online rank can call `deactivate_ranks` for one or more in-group +ranks. The Coordinator changes those members from `Active` to `Inactive`, +distributes the new membership, and waits for acknowledgements from online +ranks in the old or new active set before returning. + +Deactivation does not renumber slots or mark the target process +unhealthy. Its Agent session and data-plane links remain available; to +participate again, that process calls `join_group` and follows the normal +activation path. + +## Fault tolerance + +### Failure handling + +A failed operation reports the caller's data-plane observations to the +Coordinator; the worker itself does not make a membership decision. The two +caller-visible results are: + +- `local_success`: whether all transfers required by this operation completed + at the caller. A successful local result is valid for that caller, but says + nothing about whether the operation completed at every other rank. +- `failed_ranks_hint`: a per-operation bitmap of length `max_group_size`, + indexed by in-group rank. It records the ranks for which the caller observed + a transfer failure; a set bit is evidence, not a global conclusion that the + peer is faulty. + +Failure hints can differ across ranks. Workers submit such observations to the +Coordinator instead of changing membership locally. + +#### Reconciliation and `sync_after_failure` + +Negative evidence opens a reconciliation window in the Coordinator, allowing +reports from different ranks to arrive before a single decision is made. The +window is 30 seconds by default and must be configured to exceed the default +collective timeout. When the window closes, the Coordinator derives the +mutually connected healthy set, updates `RankState`, and distributes the +result. For groups with auto-deactivation enabled, it also changes unhealthy +active members to `Inactive`. + +`sync_after_failure` is both a reporting path and a synchronization point. Its +request piggybacks the Agent's current, unacknowledged link observations. If the +caller has just observed `local_success=false`, those observations can open or +join the Coordinator's reconciliation window. The call then waits for any +pending reconciliation and applies the group view returned by +the Coordinator. + +With `auto_deactivate_on_failure=true`, a successful return means that the +caller has applied the membership produced by reconciliation. Its local +`active_ranks` therefore reflects the Coordinator's deactivation decision, and +locally cached readiness queries such as `get_peer_state` are based on the +reconciled state. + +With auto-deactivation disabled, reconciliation does not remove members, so +the membership in the returned view may be unchanged. + +#### Failure-handling modes + +The two options control different parts of the failure path: + +- `auto_deactivate_on_failure` selects who owns failure-driven membership + changes: Mooncake PG or the upper layer framework. +- `auto_sync_on_failure` selects whether a failed collective or P2P operation + calls `sync_after_failure` automatically before it completes. It does not + control whether the Coordinator reconciles observations or automatically + changes membership. + +With auto-deactivation disabled, an unhealthy rank may remain active. An +`Offline` rank rejects new operations locally, whereas a `Synced` but unhealthy +rank may still issue them. Successful transfers provide positive link evidence +and can return a `Synced` rank to `Healthy`. + +There are three valid configurations. + +##### PG-managed, synchronized (default) + +```text +auto_deactivate_on_failure = true +auto_sync_on_failure = true +``` + +After a local transfer failure, the operation reports its observations and +automatically enters `sync_after_failure`. The Coordinator reconciles rank +state, deactivates unhealthy members, and returns the resulting view; the +caller applies that view before the failed operation completes. The framework +does not need a separate synchronization or deactivation step. + +This is the safest and simplest mode, but it deliberately puts control-plane +latency on the failure-completion path. A negative observation opens a +reconciliation window, so the failed operation may remain pending for tens of +seconds while the Coordinator reconciles reports from different ranks. + +##### PG-managed, deferred synchronization + +```text +auto_deactivate_on_failure = true +auto_sync_on_failure = false +``` + +Mooncake PG still reconciles observations and owns the failure-driven +deactivation decision. The difference is that the operation returns after its +data-plane work finishes and exposes `local_success` and `failed_ranks_hint` +without waiting for the reconciliation window. The framework may perform +other work first and call `sync_after_failure` later, before relying on the new +membership or resuming communication on the group. + +This mode is useful because `local_success` and `failed_ranks_hint` are +data-plane results, whereas reconciliation is a much slower control-plane +operation. It separates failure notification from membership synchronization +without transferring the membership decision back to the framework. + +##### Framework-managed + +```text +auto_deactivate_on_failure = false +auto_sync_on_failure = false ``` -Healthy rank responsibilities: +The failed operation returns its local evidence without changing membership. +The framework observes failures through `local_success` and +`failed_ranks_hint`, chooses the ranks to remove, and then calls +`deactivate_ranks`. -1. Reserve capacity with `max_world_size` or `extend_group_size_to()`. -2. Poll `get_peer_state(backend, ranks)` from all healthy ranks in a consistent - order. -3. Call `recover_ranks(backend, ranks)` once candidate ranks are connected. -4. Refresh higher-level components, such as Mooncake EP buffers, if they cache - transport metadata. +##### Invalid combination -Joining rank responsibilities: +```text +auto_deactivate_on_failure = false +auto_sync_on_failure = true +``` + +This combination is rejected at construction. Automatic synchronization is +meaningful only when Mooncake PG also owns failure-driven deactivation; +otherwise synchronization cannot produce an automatically updated membership +for the failed operation. + +This restriction applies only to automatic synchronization. +`sync_after_failure` may still be called manually in any mode, including when +`auto_deactivate_on_failure=false`. The call also acts as an explicit pull of +the Coordinator's latest group view, rather than relying only on pushed view +updates. The framework can therefore obtain a current view while retaining its +own deactivation policy. + +### Recovery + +A replacement process and a same-process in-place rejoin are separate ways to +bring an inactive slot back. Both reuse the scale-up flow: restored +connectivity can make a rank `Healthy`, but rejoining membership still requires +`join_group` and activation. + +#### Replacement process + +A replacement registers a new Agent session for the same global rank. The +Coordinator increments the rank epoch and invalidates the old process's link +evidence and endpoints. The replacement recreates its local communicators in +extension mode, starts at `Synced`, may perform local warmup while isolated, +and then follows the normal join and activation flow. -1. Initialize the process group with `is_extension=True`. -2. Publish local peer metadata through the backend initialization path. -3. Call `join_group(backend)` and block until healthy ranks publish extension - state. -4. Re-enter normal collectives after `join_group()` returns. +#### In-place rejoin -## Subgroup semantics +In-place rejoin applies when the process and its control-plane session remain +alive but the rank has become inactive, commonly after a transient data-plane +failure. Once connectivity recovers, the Coordinator can mark the rank +`Healthy` again. The same process calls `join_group`, drains old +work, republishes the endpoint under a fresh epoch, and +waits for activation. No process restart is required. -Mooncake PG follows PyTorch process-group ordering requirements. All processes -that participate in a parent group should call `dist.new_group()` in a consistent -order. This is especially important for elastic subgroups because healthy ranks -and joining ranks must agree on store prefixes and backend indices. +## Integration -For split-rank elastic patterns, create subgroups using the current membership -on healthy ranks and the eventual membership on joining ranks, while preserving -the same creation order. The PG elastic tests contain executable examples of this -pattern. +### Choosing an integration model -## Collective behavior +An integration combines two independent choices: how planned scaling changes +the communication group, and who owns failure-driven deactivation. -Collectives use the backend active-rank state to skip inactive ranks. The exact -implementation varies by operation and device type, but the high-level contract -is: +For planned scaling, framework-level group replacement creates another group +and switches to it, so it works with a fixed-membership CCL. Mooncake PG dynamic +membership instead keeps the same group and changes its active ranks. -- active ranks participate in the collective; -- inactive ranks are not waited on; -- if communication detects a rank failure, active-rank state can be updated; -- user code can read the current mask with `pg.get_active_ranks(backend)`. +For failure handling, `auto_deactivate_on_failure` determines ownership. When +it is `true`, Mooncake PG owns deactivation; when it is `false`, the framework +does. `auto_sync_on_failure` is separate from ownership: it only controls +whether a failed operation invokes `sync_after_failure` automatically and waits +for any pending reconciliation before completing. -The backend currently implements common collective APIs including all-reduce, -broadcast, all-gather, reduce-scatter, all-to-all, barrier, reduce, gather, -scatter, and single-tensor P2P send/recv. +The table below summarizes the four supported combinations: -## Failure and recovery boundaries +| Scaling × Failure mode | **PG-managed membership**
PG reconciles and deactivates | **Framework-managed membership**
Framework synchronizes and decides | +| ---------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| **Mooncake PG dynamic membership**
Keep the current group and change its active ranks | Mooncake PG commits requested scaling changes and handles failure-driven deactivation in the existing group. This fits a lean integration with minimal membership orchestration and no group rebuilds. | The group stays in place, while the framework synchronizes after failures and submits its own deactivation decisions. This fits applications that need direct control over group membership; it uses dynamic membership without group rebuilds, but requires an explicit failure-control path. | +| **Framework-level group replacement**
Create a standby/new group and switch to it | The framework switches groups for planned scaling, while Mooncake PG handles failure-driven deactivation in the current group. This fits an existing standby-group design with minimal failure orchestration; planned scaling still incurs the cost of group creation and switching. | The framework's control plane manages both replacement groups and failure policy. This fits frameworks that already manage group lifecycle and failure policy centrally; it offers the most flexibility but requires the most orchestration. | -Mooncake PG exposes low-level recovery primitives; higher-level systems are -responsible for policy decisions such as: +### Interfaces -- which ranks are safe to replace; -- when to stop routing traffic to a failed rank; -- how to recreate model state on a replacement process; -- when to refresh EP, scheduler, or application-level metadata; -- how to coordinate subgroup recovery. +Mooncake PG exposes a PyTorch integration and an experimental C API. -Avoid assuming that `recover_ranks()` alone reconstructs all higher-level state. -It activates the process-group communication path; the application still owns -model weights, KV-cache state, routing policy, and request scheduling. +Importing `mooncake.pg` registers two `torch.distributed` backends: -## Testing checklist for PG changes +- `mooncake-cpu` for CPU devices; +- `mooncake` for the accelerator supported by the build. -When modifying PG internals, run at least: +`MooncakeBackend` derives from `c10d::ProcessGroup`, so applications use the +usual PyTorch entry points, including `dist.init_process_group()`, +`dist.new_group()`, `dist.all_reduce()`, and `dist.batch_isend_irecv()`. +Mooncake-specific capacity, extension, and failure-handling options are passed +through `MooncakeBackendOptions`. + +PyTorch dispatches P2P operations and some collective entry points through a +`c10d::Backend` object. Each `MooncakeBackend` therefore registers a lightweight +`MooncakeBackendShim` that forwards supported operations back to its owning +`MooncakeBackend`. +See the [Python API](../api-reference/python/ep-backend.md) for +initialization examples and API details. + +Non-PyTorch integrations can use the experimental C API declared in +[`mooncake_pg.h`](https://github.com/kvcache-ai/Mooncake/blob/main/mooncake-pg/include/mooncake_pg.h). + +## Contributing + +The PG tests are the executable contracts for current behavior: + +| Test file | Contract covered | +| ---------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------ | +| `test_pg_init_functional.py` | Initialization, single rank, subgroup creation, destruction, and reinitialization | +| `test_pg_collectives.py` | Collective coverage | +| `test_pg_p2p.py` | Direct and batched P2P, ordering, multiple senders, and failure detection | +| `test_pg_elastic.py` | Automatic and manual failure handling, scale-up, process replacement, graceful leave, subgroup extension, holes, and in-place rejoin | +| `test_pg_inference_topologies.py` | TP, PP, DP, EP, and prefill/decode group layouts | +| `test_pg_inference_collectives.py` | Traffic across inference-style groups | + +Run tests from the repository root: ```bash -# CPU functional tests +# Run all PG tests +python -m unittest discover -s mooncake-pg/tests -v + +# Run CPU-only PG tests python -m unittest discover -s mooncake-pg/tests -k CPU -v -# CUDA functional tests, when GPUs are available +# Run CUDA PG tests python -m unittest discover -s mooncake-pg/tests -k CUDA -v # Collective benchmark smoke test @@ -213,9 +524,8 @@ python mooncake-pg/benchmark/pgbench.py \ --collective all_reduce --backend mooncake --device cuda -g 2 -b 8 -e 1M -f 2 ``` -Also run elastic tests for changes that touch active ranks, metadata polling, -subgroups, `extend_group_size_to()`, `get_peer_state()`, `recover_ranks()`, or -`join_group()`. +Set `MOONCAKE_PGTEST_DEVICE_FILTERS` to a comma-separated NIC/HCA list when the +test environment needs explicit device selection. ## Related documentation diff --git a/docs/source/design/mooncake-ep.md b/docs/source/design/mooncake-ep.md index cbd977e0af..aa65b3315d 100644 --- a/docs/source/design/mooncake-ep.md +++ b/docs/source/design/mooncake-ep.md @@ -16,7 +16,7 @@ Mooncake EP is designed to: - keep the Python programming model close to DeepEP low-latency mode; - use Mooncake device transports for fast intra-node and inter-node movement; - detect failed source ranks through timeout-aware kernels; -- interoperate with Mooncake Backend (PG) for bootstrap metadata exchange and +- interoperate with Mooncake PG for bootstrap metadata exchange and rank-health state. ## High-level data flow @@ -242,6 +242,6 @@ Adapt launch commands to the target environment and number of GPUs. ## Related documentation -- [Mooncake Backend (PG) design](mooncake-backend-pg.md) +- [Mooncake PG design](mooncake-backend-pg.md) - [Python API reference](../api-reference/python/ep-backend.md) - [PG/EP troubleshooting](../troubleshooting/pg-ep-troubleshooting.md) diff --git a/docs/source/troubleshooting/pg-ep-troubleshooting.md b/docs/source/troubleshooting/pg-ep-troubleshooting.md index 684e08f993..81a8ae29dc 100644 --- a/docs/source/troubleshooting/pg-ep-troubleshooting.md +++ b/docs/source/troubleshooting/pg-ep-troubleshooting.md @@ -1,7 +1,7 @@ # Mooncake PG/EP Troubleshooting This page covers common setup, import, runtime, and recovery issues for -Mooncake Backend (PG) and Mooncake EP. +Mooncake PG and Mooncake EP. ## Import fails with a PyTorch version error @@ -41,80 +41,25 @@ Fixes: 3. Make sure the Python environment used at runtime is the same one used for the build. -## `activeRanks must be int` or device mismatch +## `dist.get_world_size()` differs from the number of active ranks -Symptoms: - -```text -activeRanks must be int. -activeRanks must be on CPU. -activeRanks must be on GPU. -activeRanks must be sized to max_world_size when max_world_size is set -``` - -Causes and fixes: - -- Use `torch.int32`, not `bool`, `int64`, or floating-point types. -- For `backend="mooncake"`, put `active_ranks` on the accelerator device. -- For `backend="mooncake-cpu"`, put `active_ranks` on CPU. -- If `max_world_size` is set, allocate `active_ranks` with length - `max_world_size`, even if the initial visible `world_size` is smaller. - -Examples: - -```python -# CUDA / accelerator backend -active_ranks = torch.ones(max_world_size, dtype=torch.int32, device="cuda") - -# CPU backend -active_ranks = torch.ones(max_world_size, dtype=torch.int32) -``` - -## `dist.get_world_size()` is smaller than `max_world_size` - -This is expected. Mooncake PG distinguishes reserved capacity from visible active -membership: +Mooncake PG preserves stable in-group rank slots. Consequently: -- `max_world_size` reserves future rank slots. -- `dist.get_world_size()` returns the visible active size. -- Reserved ranks are inactive until `recover_ranks()` activates them. +- `max_group_size` is the fixed slot capacity; +- `dist.get_world_size()` is the highest active in-group rank plus one; +- holes below that extent remain visible in the rank space but are skipped by + the active mask. Use `pg.get_active_ranks(backend)` to inspect the current backend mask. -## `get_peer_state()` or `join_group()` hangs +## `join_group()` hangs or activation times out Common causes: -- Healthy ranks are not all calling `get_peer_state()` in the same order. -- The joining rank did not initialize with `is_extension=True`. -- `max_world_size` / rank numbering differs between healthy and joining ranks. -- Subgroups were created in different orders on healthy and joining processes. -- The joining process has not published peer metadata yet. -- Network device filters differ across ranks. - -Debug checklist: - -1. Confirm all ranks use the same rendezvous address and store. -2. Print rank, visible `world_size`, and `max_world_size` at initialization. -3. Confirm `active_ranks` length and values on every rank. -4. Verify healthy ranks poll the same `join_ranks` list. -5. For subgroup recovery, confirm every rank calls `dist.new_group()` in the same - order. -6. If using RDMA, set the same HCA whitelist on every rank. - -## Newly extended ranks participate too early - -After `pg.extend_group_size_to(backend, new_size)`, new ranks are reserved but -inactive. They should not participate in collectives until healthy ranks call -`pg.recover_ranks(backend, ranks)`. - -If a new rank appears to participate early, check whether: - -- `active_ranks` was initialized with `1` for future ranks without masking them - through the backend protocol; -- the application called collectives on the joining process before - `pg.join_group()` returned; -- different ranks used inconsistent `max_world_size` or rank IDs. +- No existing rank submitted `activate_ranks()` / `recover_ranks()` after the + joining process entered `join_group()`. +- The future active set is not mutually connected, so the activation proposal + remains pending until timeout. ## EP dispatch/combine timeout marks a rank inactive From 8703da6e78ac052c20f20e522e27095448e55537 Mon Sep 17 00:00:00 2001 From: Aoi Date: Mon, 10 Aug 2026 17:38:01 +0800 Subject: [PATCH 018/483] [Doc] Update PyPI package documentation (#3356) --- README.md | 1 - docs/source/getting_started/quick-start.md | 34 +++++++++------------- 2 files changed, 13 insertions(+), 22 deletions(-) diff --git a/README.md b/README.md index 5d0cb1f2ee..74276efebe 100644 --- a/README.md +++ b/README.md @@ -25,7 +25,6 @@ [![PyPI NPU](https://img.shields.io/static/v1?label=pypi&message=NPU&color=F87171)](https://pypi.org/project/mooncake-transfer-engine-npu/) [![PyPI MUSA](https://img.shields.io/static/v1?label=pypi&message=MUSA&color=F97316)](https://pypi.org/project/mooncake-transfer-engine-musa/) [![PyPI EFA CUDA 12](https://img.shields.io/static/v1?label=pypi&message=EFA%20%2B%20CUDA%2012&color=F59E0B)](https://pypi.org/project/mooncake-transfer-engine-efa/) - [![PyPI EFA CUDA 13](https://img.shields.io/static/v1?label=pypi&message=EFA%20%2B%20CUDA%2013&color=F59E0B)](https://pypi.org/project/mooncake-transfer-engine-efa-cuda13/) [![PyPI EFA Non-CUDA](https://img.shields.io/static/v1?label=pypi&message=EFA%20non-CUDA&color=F59E0B)](https://pypi.org/project/mooncake-transfer-engine-efa-non-cuda/)
diff --git a/docs/source/getting_started/quick-start.md b/docs/source/getting_started/quick-start.md index 90d17c6f30..570aa37b2c 100644 --- a/docs/source/getting_started/quick-start.md +++ b/docs/source/getting_started/quick-start.md @@ -20,27 +20,18 @@ Install the Mooncake package from PyPI. The same package provides: - Transfer Engine Python bindings and runtime components for direct `mooncake.engine.TransferEngine` usage. -**For CUDA-enabled systems:** - -- CUDA < 13.0 -```bash -pip install mooncake-transfer-engine -``` - -- CUDA >= 13.0 -```bash -pip install mooncake-transfer-engine-cuda13 -``` - -**For non-CUDA systems:** -```bash -pip install mooncake-transfer-engine-non-cuda -``` - -**For NPU systems:** -```bash -pip install mooncake-transfer-engine-npu -``` +Choose the package that matches your runtime environment. Install only one +variant in an environment. + +| Runtime environment | PyPI package | Installation command | +|---------------------|--------------|----------------------| +| NVIDIA CUDA 12.1–12.9 | [`mooncake-transfer-engine`](https://pypi.org/project/mooncake-transfer-engine/) | `pip install mooncake-transfer-engine` | +| NVIDIA CUDA 13.0/13.1 | [`mooncake-transfer-engine-cuda13`](https://pypi.org/project/mooncake-transfer-engine-cuda13/) | `pip install mooncake-transfer-engine-cuda13` | +| Non-CUDA | [`mooncake-transfer-engine-non-cuda`](https://pypi.org/project/mooncake-transfer-engine-non-cuda/) | `pip install mooncake-transfer-engine-non-cuda` | +| Ascend NPU | [`mooncake-transfer-engine-npu`](https://pypi.org/project/mooncake-transfer-engine-npu/) | `pip install mooncake-transfer-engine-npu` | +| Moore Threads MUSA | [`mooncake-transfer-engine-musa`](https://pypi.org/project/mooncake-transfer-engine-musa/) | `pip install mooncake-transfer-engine-musa` | +| AWS EFA with CUDA 12 | [`mooncake-transfer-engine-efa`](https://pypi.org/project/mooncake-transfer-engine-efa/) | `pip install mooncake-transfer-engine-efa` | +| AWS EFA without CUDA | [`mooncake-transfer-engine-efa-non-cuda`](https://pypi.org/project/mooncake-transfer-engine-efa-non-cuda/) | `pip install mooncake-transfer-engine-efa-non-cuda` | > **Important**: > - The CUDA version (`mooncake-transfer-engine`) includes Mooncake-EP and GPU topology detection, requiring CUDA 12.1+. @@ -48,6 +39,7 @@ pip install mooncake-transfer-engine-npu > ```bash > sudo apt-get update && sudo apt-get install -y libcurl4 libibverbs1 rdma-core librdmacm1 libnuma1 liburing2 > ``` +> - The EFA variants require the AWS EFA driver and libfabric at runtime. See the [EFA transport guide](../design/transfer-engine/efa_transport.md) for prerequisites and configuration. > - MLU support is currently available through source builds with `-DUSE_MLU=ON`; there is no dedicated prebuilt MLU wheel yet. > - If users encounter problems such as missing `lib*.so`, first install the corresponding system runtime libraries. If the issue persists, uninstall the package and build the binaries manually. From 64495bdfd3e215e495db714418c4b60c0f7d0fad Mon Sep 17 00:00:00 2001 From: Kamil Date: Mon, 10 Aug 2026 14:33:15 +0300 Subject: [PATCH 019/483] [Store] Allow configuring master admin/metrics HTTP bind address (#3318) * Allow configuring master admin/metrics HTTP bind address The master admin/metrics HTTP server (MasterAdminServer) always constructs coro_http_server with the 2-arg overload, so it binds to the default "0.0.0.0" (IPv4-any). On IPv6-only hosts the resolver yields an IPv4 endpoint and bind() fails, so async_start() returns an error and the master exits (restart loop). There was no way to override the bind address. Add a `metrics_host` option (gflag/config, default "0.0.0.0" for backward compatibility) plumbed through MasterConfig, MasterServiceSupervisorConfig and both the HA and non-HA construction paths into MasterAdminServer(http_port, enable_metric_reporting, http_host), which now passes it to the 3-arg coro_http_server(thread_num, port, address) overload. Set `metrics_host=::` to listen on IPv6 (dual-stack on hosts that allow it). Also include host and the async_start error message in the failure log to make bind failures diagnosable. --------- Co-authored-by: Claude Opus 4.8 (1M context) --- mooncake-store/include/master_admin_service.h | 4 +++- mooncake-store/include/master_config.h | 3 +++ .../src/ha/leadership/master_service_supervisor.cpp | 2 +- mooncake-store/src/master.cpp | 13 ++++++++++++- mooncake-store/src/master_admin_service.cpp | 10 ++++++---- 5 files changed, 25 insertions(+), 7 deletions(-) diff --git a/mooncake-store/include/master_admin_service.h b/mooncake-store/include/master_admin_service.h index 9d7dd60cb5..57b8d78068 100644 --- a/mooncake-store/include/master_admin_service.h +++ b/mooncake-store/include/master_admin_service.h @@ -23,7 +23,8 @@ class WrappedMasterService; class MasterAdminServer { public: - MasterAdminServer(uint16_t http_port, bool enable_metric_reporting); + MasterAdminServer(uint16_t http_port, bool enable_metric_reporting, + std::string http_host = "0.0.0.0"); ~MasterAdminServer(); @@ -107,6 +108,7 @@ class MasterAdminServer { void RegisterHandler(); uint16_t http_port_; + std::string http_host_; bool enable_metric_reporting_ = false; coro_http::coro_http_server http_server_; std::thread metric_report_thread_; diff --git a/mooncake-store/include/master_config.h b/mooncake-store/include/master_config.h index 2940d0e59d..4aa71b3fac 100644 --- a/mooncake-store/include/master_config.h +++ b/mooncake-store/include/master_config.h @@ -30,6 +30,7 @@ inline std::string ResolveConfiguredHABackendConnstring( struct MasterConfig { bool enable_metric_reporting; uint32_t metrics_port; + std::string metrics_host; uint32_t rpc_port; uint32_t rpc_thread_num; std::string rpc_address; @@ -182,6 +183,7 @@ class MasterServiceSupervisorConfig { // Parameters with default values (optional parameters) std::string rpc_address = "0.0.0.0"; + std::string metrics_host = "0.0.0.0"; std::chrono::steady_clock::duration rpc_conn_timeout = std::chrono::seconds( 0); // Client connection timeout. 0 = no timeout (infinite) bool rpc_enable_tcp_no_delay = true; @@ -269,6 +271,7 @@ class MasterServiceSupervisorConfig { // Set required parameters using RequiredParam enable_metric_reporting = config.enable_metric_reporting; metrics_port = static_cast(config.metrics_port); + metrics_host = config.metrics_host; default_kv_lease_ttl = config.default_kv_lease_ttl; default_kv_soft_pin_ttl = config.default_kv_soft_pin_ttl; allow_evict_soft_pinned_objects = diff --git a/mooncake-store/src/ha/leadership/master_service_supervisor.cpp b/mooncake-store/src/ha/leadership/master_service_supervisor.cpp index a60995a353..95b51a72cb 100644 --- a/mooncake-store/src/ha/leadership/master_service_supervisor.cpp +++ b/mooncake-store/src/ha/leadership/master_service_supervisor.cpp @@ -516,7 +516,7 @@ int MasterServiceSupervisor::Start() { mooncake::MasterAdminServer admin_server( static_cast(config_.metrics_port), - config_.enable_metric_reporting); + config_.enable_metric_reporting, config_.metrics_host); if (!admin_server.Start()) { LOG(ERROR) << "Failed to start master admin server, metrics_port=" << config_.metrics_port; diff --git a/mooncake-store/src/master.cpp b/mooncake-store/src/master.cpp index e4e9dec07c..c9fa9b5a94 100644 --- a/mooncake-store/src/master.cpp +++ b/mooncake-store/src/master.cpp @@ -118,6 +118,9 @@ DEFINE_int32( "Maximum number of threads to use (deprecated, use rpc_thread_num)"); DEFINE_bool(enable_metric_reporting, true, "Enable periodic metric reporting"); DEFINE_int32(metrics_port, 9003, "Port for HTTP metrics server to listen on"); +DEFINE_string(metrics_host, "0.0.0.0", + "Address for the HTTP metrics/admin server to listen on. " + "Use \"::\" to listen on IPv6 (and IPv4 on dual-stack hosts)"); DEFINE_string(default_kv_lease_ttl, kDefaultKvLeaseTtlFlagValue, "Default lease time for kv objects. Supports raw milliseconds " "or duration strings with ms, s, m, or h suffixes"); @@ -440,6 +443,8 @@ void InitMasterConf(const mooncake::DefaultConfig& default_config, FLAGS_enable_metric_reporting); default_config.GetUInt32("metrics_port", &master_config.metrics_port, FLAGS_metrics_port); + default_config.GetString("metrics_host", &master_config.metrics_host, + FLAGS_metrics_host); default_config.GetUInt32("rpc_port", &master_config.rpc_port, FLAGS_rpc_port); default_config.GetUInt32("rpc_thread_num", &master_config.rpc_thread_num, @@ -780,6 +785,11 @@ void LoadConfigFromCmdline(mooncake::MasterConfig& master_config, !conf_set) { master_config.metrics_port = FLAGS_metrics_port; } + if ((google::GetCommandLineFlagInfo("metrics_host", &info) && + !info.is_default) || + !conf_set) { + master_config.metrics_host = FLAGS_metrics_host; + } if ((google::GetCommandLineFlagInfo("default_kv_lease_ttl", &info) && !info.is_default) || !conf_set) { @@ -1421,6 +1431,7 @@ int main(int argc, char* argv[]) { << ", max_threads=" << master_config.rpc_thread_num << ", enable_metric_reporting=" << master_config.enable_metric_reporting << ", metrics_port=" << master_config.metrics_port + << ", metrics_host=" << master_config.metrics_host << ", default_kv_lease_ttl=" << master_config.default_kv_lease_ttl << ", default_kv_soft_pin_ttl=" << master_config.default_kv_soft_pin_ttl << ", allow_evict_soft_pinned_objects=" @@ -1541,7 +1552,7 @@ int main(int argc, char* argv[]) { metadata_server_ptr, http_metadata_remote_url); mooncake::MasterAdminServer admin_server( static_cast(master_config.metrics_port), - master_config.enable_metric_reporting); + master_config.enable_metric_reporting, master_config.metrics_host); if (!admin_server.Start()) { LOG(ERROR) << "Failed to start master admin server"; return 1; diff --git a/mooncake-store/src/master_admin_service.cpp b/mooncake-store/src/master_admin_service.cpp index 950f02fd60..2b4d3fbf77 100644 --- a/mooncake-store/src/master_admin_service.cpp +++ b/mooncake-store/src/master_admin_service.cpp @@ -252,10 +252,12 @@ tl::expected ParseQuotaPolicyBody( } // namespace MasterAdminServer::MasterAdminServer(uint16_t http_port, - bool enable_metric_reporting) + bool enable_metric_reporting, + std::string http_host) : http_port_(http_port), + http_host_(std::move(http_host)), enable_metric_reporting_(enable_metric_reporting), - http_server_(4, http_port) {} + http_server_(4, http_port, http_host_) {} MasterAdminServer::~MasterAdminServer() { Stop(); } @@ -268,8 +270,8 @@ bool MasterAdminServer::Start() { auto ec = http_server_.async_start(); if (ec.hasResult()) { - LOG(ERROR) << "Failed to start master admin server on port " - << http_port_; + LOG(ERROR) << "Failed to start master admin server on " << http_host_ + << ":" << http_port_ << ": " << ec.value().message(); return false; } From 228d49ce6c98824b1b1c4e632615ef85aee8567a Mon Sep 17 00:00:00 2001 From: Aoi Date: Mon, 10 Aug 2026 22:36:44 +0800 Subject: [PATCH 020/483] [CI/Build] Fix nightly CUDA stub runtime (#3346) * [CI/Build] Fix nightly CUDA stub runtime * [CI/Build] Fix nightly CUDA consumer linkage --- .github/workflows/nightly.yml | 22 ++++++++++--- scripts/ci/run_store_go_integration.sh | 45 +++++++++++--------------- 2 files changed, 37 insertions(+), 30 deletions(-) diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index ec31a8082c..6976f4c8d5 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -248,10 +248,26 @@ jobs: -DSTORE_USE_ETCD=ON -DCMAKE_BUILD_TYPE=Release \ -DBUILD_UNIT_TESTS=ON -DENABLE_SCCACHE=ON + - name: Configure CUDA driver runtime + run: | + cuda_driver_library=$(sed -n \ + 's/^CUDA_cuda_driver_LIBRARY:FILEPATH=//p' build/CMakeCache.txt) + if [ -z "$cuda_driver_library" ] || [ ! -f "$cuda_driver_library" ]; then + echo "::error::CMake did not resolve the CUDA driver library" + exit 1 + fi + + cuda_driver_dir=$(dirname "$cuda_driver_library") + if [ ! -e "$cuda_driver_dir/libcuda.so.1" ]; then + sudo ln -s "$(basename "$cuda_driver_library")" \ + "$cuda_driver_dir/libcuda.so.1" + fi + echo "LIBRARY_PATH=$cuda_driver_dir:${LIBRARY_PATH:-}" >> "$GITHUB_ENV" + echo "LD_LIBRARY_PATH=$cuda_driver_dir:${LD_LIBRARY_PATH:-}" >> "$GITHUB_ENV" + - name: Build project run: | cd build - export LIBRARY_PATH=/usr/local/cuda/lib64/stubs:${LIBRARY_PATH:-} cmake --build . -j$(nproc) sudo -E cmake --install . @@ -259,7 +275,6 @@ jobs: run: | mkdir -p build/mooncake-transfer-engine/nvlink-allocator cd mooncake-transfer-engine/nvlink-allocator - export LIBRARY_PATH=/usr/local/cuda/lib64/stubs:$LIBRARY_PATH bash build.sh ../../build/mooncake-transfer-engine/nvlink-allocator/ - name: Start Metadata Server @@ -272,7 +287,7 @@ jobs: - name: Run CTest unit tests run: | cd build - export LD_LIBRARY_PATH=$LD_LIBRARY_PATH:/usr/local/lib + export LD_LIBRARY_PATH=${LD_LIBRARY_PATH:-}:/usr/local/lib MC_METADATA_SERVER=http://127.0.0.1:8080/metadata \ DEFAULT_KV_LEASE_TTL=500 \ ctest --parallel $(nproc) --output-on-failure @@ -286,7 +301,6 @@ jobs: - name: Run Go store binding integration tests env: MOONCAKE_STORE_CLUSTER_ID: nightly_go_cluster - MOONCAKE_STORE_GO_LINK_COMMON: "0" MOONCAKE_STORE_GO_SANITIZED: "0" run: ./scripts/ci/run_store_go_integration.sh diff --git a/scripts/ci/run_store_go_integration.sh b/scripts/ci/run_store_go_integration.sh index 482f93cf6c..0ffe8429a3 100755 --- a/scripts/ci/run_store_go_integration.sh +++ b/scripts/ci/run_store_go_integration.sh @@ -14,15 +14,6 @@ case "${MOONCAKE_STORE_GO_SANITIZED:-0}" in ;; esac -case "${MOONCAKE_STORE_GO_LINK_COMMON:-0}" in - 0) link_common=false ;; - 1) link_common=true ;; - *) - echo "MOONCAKE_STORE_GO_LINK_COMMON must be 0 or 1" >&2 - exit 2 - ;; -esac - "$GITHUB_WORKSPACE/build/mooncake-store/src/mooncake_master" \ --eviction_high_watermark_ratio=0.95 \ --cluster_id="$MOONCAKE_STORE_CLUSTER_ID" \ @@ -31,7 +22,7 @@ master_pid=$! sleep 3 cd "$GITHUB_WORKSPACE/mooncake-store/go" -export LD_LIBRARY_PATH="$GITHUB_WORKSPACE/build/mooncake-common:$GITHUB_WORKSPACE/build/mooncake-store/src:$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src:$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src/common/base:$GITHUB_WORKSPACE/build/mooncake-common/etcd" +export LD_LIBRARY_PATH="$GITHUB_WORKSPACE/build/mooncake-common:$GITHUB_WORKSPACE/build/mooncake-store/src:$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src:$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src/common/base:$GITHUB_WORKSPACE/build/mooncake-common/etcd:${LD_LIBRARY_PATH:-}" export CGO_ENABLED=1 export CGO_CFLAGS="-I$GITHUB_WORKSPACE/mooncake-store/include -I$GITHUB_WORKSPACE/mooncake-transfer-engine/include" @@ -41,23 +32,13 @@ linker_flags=( "-L$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src" "-L$GITHUB_WORKSPACE/build/mooncake-transfer-engine/src/common/base" "-L$GITHUB_WORKSPACE/build/mooncake-common" + "-L$GITHUB_WORKSPACE/build/mooncake-common/src" + "-L$GITHUB_WORKSPACE/build/mooncake-common/etcd" + -Wl,--start-group + -lmooncake_store -lcachelib_memory_allocator -ltransfer_engine -lbase + -lmooncake_common + -Wl,--end-group ) -if $link_common; then - linker_flags+=("-L$GITHUB_WORKSPACE/build/mooncake-common/src") -fi -linker_flags+=("-L$GITHUB_WORKSPACE/build/mooncake-common/etcd") -if $link_common; then - linker_flags+=( - -Wl,--start-group - -lmooncake_store -lcachelib_memory_allocator -ltransfer_engine -lbase - -lmooncake_common - -Wl,--end-group - ) -else - linker_flags+=( - -lmooncake_store -lcachelib_memory_allocator -ltransfer_engine -lbase - ) -fi linker_flags+=( -lasio -letcd_wrapper -lstdc++ -lnuma -lglog -lgflags -libverbs -lmlx5 -ljsoncpp -lzstd -lcurl -luring @@ -76,6 +57,18 @@ export CGO_LDFLAGS="${linker_flags[*]}" if [ -d /usr/local/cuda/lib64 ]; then export CGO_LDFLAGS="$CGO_LDFLAGS -L/usr/local/cuda/lib64 -lcudart" fi +# USE_CUDA links transfer_engine against the CUDA Driver API in addition to the +# runtime API. Reuse the exact driver library that CMake selected. +if grep -q '^USE_CUDA:BOOL=ON$' "$GITHUB_WORKSPACE/build/CMakeCache.txt"; then + cuda_driver_library=$(sed -n \ + 's/^CUDA_cuda_driver_LIBRARY:FILEPATH=//p' \ + "$GITHUB_WORKSPACE/build/CMakeCache.txt") + if [ -z "$cuda_driver_library" ] || [ ! -f "$cuda_driver_library" ]; then + echo "CMake did not resolve the CUDA driver library" >&2 + exit 1 + fi + export CGO_LDFLAGS="$CGO_LDFLAGS -L$(dirname "$cuda_driver_library") -lcuda" +fi # The KV events publisher is optional and linked when libzmq is installed. if ldconfig -p 2>/dev/null | grep -q libzmq; then export CGO_LDFLAGS="$CGO_LDFLAGS -lzmq" From a2933849417259e562fc9cbad63c618d464dfb2d Mon Sep 17 00:00:00 2001 From: Posedge_Lin Date: Tue, 11 Aug 2026 10:49:49 +0800 Subject: [PATCH 021/483] [Transfer Engine] Schedule TCP transfers on bounded per-peer connection lanes (#2974) * transfer-engine: schedule TCP transfers on bounded per-peer connection lanes * test: avoid requiring lane retry timer to fire * ci: rerun checks * fix: complete TCP task-group tail during shutdown * fix: harden TCP lane reconnect scheduling * style: format TCP lane reconnect condition --- .../transport/tcp_transport/tcp_transport.h | 296 ++- .../transport/tcp_transport/CMakeLists.txt | 7 + .../transport/tcp_transport/tcp_transport.cpp | 1372 +++------- .../tcp_transport/tcp_transport_lane_impl.h | 1390 ++++++++++ .../tcp_transport_session_impl.h | 815 ++++++ mooncake-transfer-engine/tests/CMakeLists.txt | 2 + .../tests/tcp_write_visibility_test.cpp | 2250 ++++++++++++++++- 7 files changed, 5066 insertions(+), 1066 deletions(-) create mode 100644 mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_lane_impl.h create mode 100644 mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_session_impl.h diff --git a/mooncake-transfer-engine/include/transport/tcp_transport/tcp_transport.h b/mooncake-transfer-engine/include/transport/tcp_transport/tcp_transport.h index 88b51c9169..9a499b07f4 100644 --- a/mooncake-transfer-engine/include/transport/tcp_transport/tcp_transport.h +++ b/mooncake-transfer-engine/include/transport/tcp_transport/tcp_transport.h @@ -19,19 +19,25 @@ #include #include +#include #include +#include #include #include #include #include #include +#include #include +#include +#include #include -#include +#include #include -#include +#include #include +#include #include "transfer_metadata.h" #include "transport/transport.h" @@ -39,26 +45,10 @@ namespace mooncake { class TransferMetadata; +struct ClientSession; class TcpContext; class TcpTransport; -// Pooled connection for reusing client-side connections -struct PooledConnection { - std::shared_ptr socket; - std::string host; - uint16_t port; - std::chrono::steady_clock::time_point last_used; - bool in_use; - - PooledConnection(std::shared_ptr s, - const std::string &h, uint16_t p) - : socket(std::move(s)), - host(h), - port(p), - last_used(std::chrono::steady_clock::now()), - in_use(true) {} -}; - class TcpTransport : public Transport { public: using BufferDesc = TransferMetadata::BufferDesc; @@ -122,10 +112,9 @@ class TcpTransport : public Transport { TcpContext *context_; std::atomic_bool running_; std::thread thread_; - bool enable_connection_pool_ = - false; // Use MC_TCP_ENABLE_CONNECTION_POOL=1 to enable connection pool + bool enable_connection_pool_ = false; - // Client-side connection pool + // Client-side bounded work queues and fixed connection lanes. struct ConnectionKey { std::string host; uint16_t port; @@ -142,23 +131,256 @@ class TcpTransport : public Transport { } }; - std::unordered_map>, - ConnectionKeyHash> - connection_pool_; - std::mutex pool_mutex_; + enum class WorkFailureReason { + QUEUE_FULL, + QUEUE_TIMEOUT, + RUNTIME_UNAVAILABLE, + CONNECT_FAILED, + SESSION_FAILED, + SHUTDOWN, + }; + + struct FailureCounters { + std::atomic queue_full{0}; + std::atomic queue_timeout{0}; + std::atomic connect_failed{0}; + std::atomic runtime_unavailable{0}; + std::atomic session_failed{0}; + std::atomic shutdown{0}; + }; + + struct TcpWorkItem { + Slice *slice = nullptr; + bool use_v2 = false; + std::function continuation; + std::chrono::steady_clock::time_point admission_deadline; + + TcpWorkItem() = default; + TcpWorkItem(Slice *slice_arg, bool use_v2_arg, + std::function continuation_arg = nullptr) + : slice(slice_arg), + use_v2(use_v2_arg), + continuation(std::move(continuation_arg)) {} + TcpWorkItem(TcpWorkItem &&other) noexcept + : slice(other.slice), + use_v2(other.use_v2), + continuation(std::move(other.continuation)), + admission_deadline(other.admission_deadline) { + other.slice = nullptr; + } + TcpWorkItem &operator=(TcpWorkItem &&) = delete; + TcpWorkItem(const TcpWorkItem &) = delete; + TcpWorkItem &operator=(const TcpWorkItem &) = delete; + }; + + struct TerminalAction { + TcpWorkItem work; + TransferStatusEnum status; + bool connection_clean; + + TerminalAction(TcpWorkItem &&work_arg, TransferStatusEnum status_arg, + bool clean_arg) + : work(std::move(work_arg)), + status(status_arg), + connection_clean(clean_arg) {} + TerminalAction(TerminalAction &&) noexcept = default; + TerminalAction &operator=(TerminalAction &&) = delete; + TerminalAction(const TerminalAction &) = delete; + TerminalAction &operator=(const TerminalAction &) = delete; + }; + + static_assert(!std::is_copy_constructible::value, + "TcpWorkItem must remain move-only"); + static_assert(!std::is_copy_assignable::value, + "TcpWorkItem must remain move-only"); + static_assert(!std::is_copy_constructible::value, + "TerminalAction must remain move-only"); + static_assert(!std::is_copy_assignable::value, + "TerminalAction must remain move-only"); + + enum class GroupState { OPEN, CLOSING, CLOSED }; + enum class LaneState { + DISCONNECTED, + CONNECTING, + IDLE, + BUSY, + COMPLETING, + CLOSING, + CLOSED, + }; + enum class LaneConnectStage { NONE, RESOLVING, CONNECTING }; + + struct ConnectionLaneRuntime { + explicit ConnectionLaneRuntime(asio::io_context &io_context) + : executor(io_context.get_executor()) {} + + asio::io_context::executor_type executor; + }; + + struct PeerConnectionGroup; + + struct ConnectionLane { + ConnectionLane(size_t lane_id_arg, + const std::shared_ptr &group_arg) + : lane_id(lane_id_arg), group(group_arg) {} + + size_t lane_id; + std::weak_ptr group; + LaneState state = LaneState::DISCONNECTED; + LaneConnectStage connect_stage = LaneConnectStage::NONE; + uint64_t operation_epoch = 0; + uint64_t last_connect_round = 0; + std::shared_ptr resolver; + std::shared_ptr socket; + std::shared_ptr session; + std::optional current; + }; + + struct PeerConnectionGroup { + PeerConnectionGroup(ConnectionKey key_arg, + const asio::io_context::executor_type &executor_arg, + size_t queue_capacity_arg, + size_t pending_admission_capacity_arg, + std::chrono::milliseconds admission_timeout_arg, + std::shared_ptr counters_arg) + : key(std::move(key_arg)), + executor(executor_arg), + queue_capacity(queue_capacity_arg), + pending_admission_capacity(pending_admission_capacity_arg), + admission_timeout(admission_timeout_arg), + failure_counters(std::move(counters_arg)) {} + + std::mutex mutex; + GroupState state = GroupState::OPEN; + ConnectionKey key; + asio::io_context::executor_type executor; + size_t queue_capacity; + std::deque queue; + std::deque pending_admissions; + size_t pending_admission_capacity; + std::chrono::milliseconds admission_timeout; + std::shared_ptr admission_timer; + uint64_t admission_epoch = 0; + uint64_t queued_bytes = 0; + bool queued_bytes_saturated = false; + std::vector> lanes; + bool pump_scheduled = false; + uint64_t pump_epoch = 0; + uint64_t connect_round = 1; + size_t probes_in_flight = 0; + bool connect_round_had_success = false; + std::chrono::steady_clock::time_point next_probe_not_before{}; + std::shared_ptr retry_timer; + uint64_t retry_epoch = 0; + uint64_t connect_failure_log_count = 0; + std::shared_ptr failure_counters; + }; + + struct ConnectionLaneState { + std::mutex mutex; + bool shutting_down = false; + size_t lanes_per_peer = 4; + size_t max_queued_transfers_per_peer = 1024; + size_t max_pending_admissions_per_peer = 1024; + std::chrono::milliseconds admission_timeout{1000}; + std::unordered_map, + ConnectionKeyHash> + groups; + std::weak_ptr runtime; + std::shared_ptr failure_counters = + std::make_shared(); + }; + + std::shared_ptr lane_runtime_; + std::shared_ptr lane_state_; + + // TODO(#2930): The queue item bound does not bound queued source bytes, + // and a connected payload without a progress deadline can still stall a + // lane indefinitely. Peer-generation recovery is also a separate phase. std::shared_ptr getConnection( - const std::string &host, uint16_t port, bool use_pool); - void returnConnection(const std::string &host, uint16_t port, - std::shared_ptr socket); - // Close `socket` and drop it from the pool: used for requests that did - // not terminate in a well-defined protocol state (#2086). - void discardConnection(const std::string &host, uint16_t port, - std::shared_ptr socket); - void cleanupIdleConnections(); - - static constexpr std::chrono::seconds kConnectionIdleTimeout{60}; + const std::string &host, uint16_t port); + void enqueuePooledTransfer(const ConnectionKey &key, TcpWorkItem work); + static uint64_t requestGroupPumpLocked(PeerConnectionGroup &group); + static void postGroupPump(const std::shared_ptr &group, + uint64_t pump_epoch); + static void runGroupPump(const std::shared_ptr &group, + uint64_t pump_epoch); + static void startLaneConnect( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch); + static void handleLaneResolved( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch, + asio::error_code ec, asio::ip::tcp::resolver::results_type results); + static void handleLaneConnected( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch, + asio::error_code ec); + static void handleLaneConnectFailure( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch, + const std::string &error); + static bool armRetryTimerLocked( + const std::shared_ptr &group); + static void handleRetryTimer( + const std::shared_ptr &group, + const std::shared_ptr &timer, uint64_t retry_epoch, + asio::error_code ec); + static size_t expirePendingAdmissionsLocked( + PeerConnectionGroup &group, std::chrono::steady_clock::time_point now, + std::deque &expired); + static size_t promotePendingAdmissionsLocked(PeerConnectionGroup &group); + static void refreshAdmissionTimerLocked( + const std::shared_ptr &group, + std::deque &runtime_failed, + std::shared_ptr &timer_to_cancel, + bool &timer_armed); + static void handleAdmissionTimer( + const std::shared_ptr &group, + const std::shared_ptr &timer, + uint64_t admission_epoch, asio::error_code ec); + static void startLaneSession( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch); + static void handleLaneTerminal( + const std::shared_ptr &group, + const std::shared_ptr &lane, uint64_t epoch, + TransferStatusEnum status, bool connection_clean) noexcept; + static void completeTerminalAction(TerminalAction action) noexcept; + static void failWorkItem( + TcpWorkItem work, WorkFailureReason reason, + const std::shared_ptr &counters) noexcept; + static void failWorkItems( + std::deque work, WorkFailureReason reason, + const std::shared_ptr &counters) noexcept; + static uint64_t recordWorkFailure( + WorkFailureReason reason, + const std::shared_ptr &counters) noexcept; + static bool hasUsableLaneLocked(const PeerConnectionGroup &group); + static bool hasDisconnectedLaneLocked(const PeerConnectionGroup &group); + static bool hasUntriedDisconnectedLaneLocked( + const PeerConnectionGroup &group); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + static size_t activeSocketCountLocked(const PeerConnectionGroup &group); +#endif + static void beginConnectRoundLocked(PeerConnectionGroup &group); + static void enterReconnectCooldownLocked(PeerConnectionGroup &group); + static void addQueuedBytesLocked(PeerConnectionGroup &group, + uint64_t length); + static void removeQueuedBytesLocked(PeerConnectionGroup &group, + uint64_t length); + static void clearQueuedBytesLocked(PeerConnectionGroup &group); + static void closeSocketNoThrow( + const std::shared_ptr &socket) noexcept; + void shutdownConnectionLanes(); + void startTransferWithSocket( + TcpWorkItem work, + std::shared_ptr socket) noexcept; + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + friend bool tcpTransportLaneTypesAreMoveOnlyForTest() noexcept; +#endif }; } // namespace mooncake diff --git a/mooncake-transfer-engine/src/transport/tcp_transport/CMakeLists.txt b/mooncake-transfer-engine/src/transport/tcp_transport/CMakeLists.txt index f2c21ea006..0c6d72bcaa 100644 --- a/mooncake-transfer-engine/src/transport/tcp_transport/CMakeLists.txt +++ b/mooncake-transfer-engine/src/transport/tcp_transport/CMakeLists.txt @@ -1,8 +1,15 @@ file(GLOB TCP_SOURCES "*.cpp") +file(GLOB TCP_IMPL_HEADERS "*_impl.h") if (USE_HIP) + list(APPEND TCP_SOURCES ${TCP_IMPL_HEADERS}) hipify_files(TCP_SOURCES) endif() add_library(tcp_transport OBJECT ${TCP_SOURCES}) target_link_libraries(tcp_transport PRIVATE JsonCpp::JsonCpp yalantinglibs::yalantinglibs) + +if(BUILD_UNIT_TESTS) + target_compile_definitions(tcp_transport + PRIVATE MOONCAKE_TCP_TRANSPORT_TEST_HOOKS) +endif() diff --git a/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport.cpp b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport.cpp index c2aa41bbcd..38f7af0010 100644 --- a/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport.cpp +++ b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport.cpp @@ -17,18 +17,25 @@ #include #include #include +#include #include #include +#include #include #include #include +#include #include #include #include +#include +#include +#include #include +#include #include -#include +#include #include "common.h" #include "transfer_engine.h" @@ -40,807 +47,199 @@ namespace mooncake { using tcpsocket = asio::ip::tcp::socket; -static size_t getChunkSize() { - static const size_t val = [] { - const char* env = std::getenv("MC_TCP_SLICE_SIZE"); - if (env) { - try { - size_t v = std::stoull(env); - if (v > 0) return v; - LOG(WARNING) - << "Ignore non-positive MC_TCP_SLICE_SIZE value: " << env - << ", using default 65536"; - } catch (const std::exception& e) { - // A non-numeric or out-of-range value makes std::stoull throw; - // fall through to the default instead of letting the exception - // propagate out of this static initializer and abort the - // transfer that first reads the chunk size. - LOG(WARNING) - << "Invalid MC_TCP_SLICE_SIZE value: " << env - << ". Error: " << e.what() << ", using default 65536"; - } - } - return size_t(65536); // 64KB default - }(); - return val; -} -struct SessionHeader { - uint64_t size; - uint64_t addr; - uint8_t opcode; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +namespace { +using LaneConnectHandlerHook = void (*)() noexcept; +using LaneConnectFailureInjectionHook = bool (*)(size_t) noexcept; +using LaneRetryHandlerHook = void (*)() noexcept; +using LaneAdmissionHandlerHook = void (*)() noexcept; +using LaneObserverHook = void (*)(int, size_t, uint64_t, size_t, bool) noexcept; +using LaneFailureReasonHook = void (*)(int) noexcept; + +std::mutex lane_test_hook_mutex; +LaneConnectHandlerHook lane_connect_handler_hook = nullptr; +LaneConnectFailureInjectionHook lane_connect_failure_injection_hook = nullptr; +LaneRetryHandlerHook lane_retry_handler_hook = nullptr; +LaneAdmissionHandlerHook lane_admission_handler_hook = nullptr; +LaneObserverHook lane_observer_hook = nullptr; +LaneFailureReasonHook lane_failure_reason_hook = nullptr; + +enum LaneTestEvent { + kLaneQueueAdmitted = 1, + kLaneQueueRejected = 2, + kLaneConnecting = 3, + kLaneBusy = 4, + kLaneTerminal = 5, + kLaneShutdownClean = 6, + kLaneLateHandler = 7, + kLaneRetryArmed = 8, + kLaneRetryFired = 9, + kLaneRetryLate = 10, + kLaneCooldownStarted = 11, + kLaneAdmissionPending = 12, + kLaneAdmissionPromoted = 13, + kLaneAdmissionTimerArmed = 14, + kLaneAdmissionTimerFired = 15, + kLaneAdmissionTimerLate = 16, + kLaneAdmissionHardRejected = 17, }; -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) -static bool isCudaMemory(void* addr) { - cudaPointerAttributes attributes; - auto status = cudaPointerGetAttributes(&attributes, addr); - if (status != cudaSuccess) return false; - return attributes.type == cudaMemoryTypeDevice; -} - -// Returns the CUDA device ordinal if addr is device memory, or -1 otherwise. -// Callers must call cudaSetDevice before any cudaMemcpy to avoid implicit -// GPU 0 context creation. -static int getCudaDeviceId(void* addr) { - cudaPointerAttributes attributes; - auto status = cudaPointerGetAttributes(&attributes, addr); - if (status != cudaSuccess) return -1; - if (attributes.type == cudaMemoryTypeDevice) return attributes.device; - return -1; -} - -#ifdef USE_MACA -static cudaError_t copyTcpCudaMemory(void* dst, const void* src, size_t size) { - cudaStream_t stream; - cudaError_t status = - cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking); - if (status != cudaSuccess) return status; - - status = cudaMemcpyAsync(dst, src, size, cudaMemcpyDefault, stream); - if (status == cudaSuccess) { - status = cudaStreamSynchronize(stream); +void invokeLaneConnectHandlerHook() noexcept { + LaneConnectHandlerHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_connect_handler_hook; } - - cudaError_t destroy_status = cudaStreamDestroy(stream); - return status == cudaSuccess ? destroy_status : status; + if (hook) hook(); } -#endif -#endif -// Forward declaration -class TcpTransport; - -using ValidateAddrFn = std::function; - -// --- Acknowledged framing (protocol v2, #2086) ------------------------------ -// v1 framing gives the initiator no channel to learn whether the receiver -// applied (or even accepted) a WRITE: COMPLETED fires when the final chunk -// enters the initiator's kernel socket buffer, while megabytes may still be -// in flight toward destination memory, and a rejected request is silently -// "successful". v2 requests set the high bit of the opcode; the server then -// (a) prefixes every READ response with an 8-byte status frame and (b) sends -// an 8-byte status frame for WRITE only after the final chunk has been -// applied to destination memory. Initiators enable v2 only when the target -// segment advertises tcp_proto_version >= 2, so old servers never see -// flagged opcodes and old initiators keep receiving v1 framing. -static constexpr uint8_t kOpcodeV2Flag = 0x80; -// Status frames carry a magic in the high 32 bits so that a v2 initiator -// which reaches a v1 server through a stale descriptor (v1 treats unknown -// opcodes as READ and immediately streams payload bytes) fails fast on the -// first frame instead of misinterpreting the stream. Residual risk: payload -// bytes that happen to equal a valid frame (2^-64 per request, data -// dependent) are indistinguishable in-band; eliminating that would need a -// nonce/checksum handshake, which this deliberately avoids. -static constexpr uint64_t kStatusMagic = 0x4D435456ull << 32; // "MCTV" -static constexpr uint64_t kStatusOk = kStatusMagic | 0; -static constexpr uint64_t kStatusAddrRejected = kStatusMagic | 1; -static inline bool statusFrameValid(uint64_t frame) { - return (frame & 0xFFFFFFFF00000000ull) == kStatusMagic; +bool invokeLaneConnectFailureInjectionHook(size_t lane_id) noexcept { + LaneConnectFailureInjectionHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_connect_failure_injection_hook; + } + return hook && hook(lane_id); } -// Operational escape hatch: MC_TCP_PROTO=1 forces initiators to speak the -// legacy unacknowledged framing even to v2-capable servers. Also used by -// tests to cover the mixed-version matrix in one process. -static bool forceLegacyTcpProto() { - // Read per call (startTransfer already does metadata lookups; getenv is - // noise) so tests can cover both protocol modes in one process. - const char* env = std::getenv("MC_TCP_PROTO"); - return env && env[0] == '1' && env[1] == '\0'; +void invokeLaneRetryHandlerHook() noexcept { + LaneRetryHandlerHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_retry_handler_hook; + } + if (hook) hook(); } -// Server-side session: handles transfer requests on a persistent connection. -// The session owns the socket; ending the callback chain without rearming -// (start()/next handler) drops the last reference and closes the connection. -struct ServerSession : public std::enable_shared_from_this { - explicit ServerSession(std::shared_ptr socket, - ValidateAddrFn validate_addr) - : socket_(std::move(socket)), - validate_addr_(std::move(validate_addr)) {} - - std::shared_ptr socket_; - ValidateAddrFn validate_addr_; - SessionHeader header_; - uint64_t total_transferred_bytes_; - char* local_buffer_; - bool v2_ = false; - uint64_t status_frame_; - - void start() { - total_transferred_bytes_ = 0; - readHeader(); +void invokeLaneAdmissionHandlerHook() noexcept { + LaneAdmissionHandlerHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_admission_handler_hook; } + if (hook) hook(); +} - private: - // Send an 8-byte status frame, then run `next` (or end the session — - // closing the connection — when `next` is empty or the send fails). - void sendStatus(uint64_t status, std::function next) { - status_frame_ = htole64(status); - auto self(shared_from_this()); - asio::async_write(*socket_, - asio::buffer(&status_frame_, sizeof(status_frame_)), - [this, self, next = std::move(next)]( - const asio::error_code& ec, std::size_t) { - if (ec) - return; // connection closes with the session - if (next) next(); - }); +void invokeLaneObserverHook(int event, size_t queue_depth, + uint64_t queued_bytes, size_t active_sockets, + bool lane_has_current) noexcept { + LaneObserverHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_observer_hook; } + if (hook) + hook(event, queue_depth, queued_bytes, active_sockets, + lane_has_current); +} - void readHeader() { - auto self(shared_from_this()); - asio::async_read( - *socket_, asio::buffer(&header_, sizeof(SessionHeader)), - [this, self](const asio::error_code& ec, std::size_t len) { - if (ec || len != sizeof(SessionHeader)) { - if (ec.value() != asio::error::eof) { - LOG(WARNING) - << "ServerSession::readHeader failed. Error: " - << ec.message() << " (value: " << ec.value() << ")" - << ", bytes read: " << len; - } - return; - } - - v2_ = (header_.opcode & kOpcodeV2Flag) != 0; - const uint8_t opcode = header_.opcode & ~kOpcodeV2Flag; - local_buffer_ = (char*)(le64toh(header_.addr)); - uint64_t size = le64toh(header_.size); - if (validate_addr_ && - !validate_addr_((uint64_t)local_buffer_, size)) { - LOG(ERROR) << "ServerSession: remote-supplied address 0x" - << std::hex << (uint64_t)local_buffer_ - << std::dec << " with size " << size - << " is not within any registered buffer"; - // v2 initiators learn of the rejection; v1 initiators - // only see the connection close (and, for small WRITEs, - // may have already reported success — the defect v2 - // exists to fix). - if (v2_) sendStatus(kStatusAddrRejected, nullptr); - return; - } - if (opcode == (uint8_t)TransferRequest::WRITE) { - readBody(); - } else if (v2_) { - // READ, v2: status frame precedes the data. - sendStatus(kStatusOk, [this] { writeBody(); }); - } else { - writeBody(); - } - }); +void invokeLaneFailureReasonHook(int reason) noexcept { + LaneFailureReasonHook hook; + { + std::lock_guard lock(lane_test_hook_mutex); + hook = lane_failure_reason_hook; } + if (hook) hook(reason); +} +} // namespace - void writeBody() { - auto self(shared_from_this()); - uint64_t size = le64toh(header_.size); - char* addr = local_buffer_; - - size_t buffer_size = - std::min(getChunkSize(), size - total_transferred_bytes_); - if (buffer_size == 0) { - // Transfer complete, wait for next request on this connection - start(); - return; - } - - char* dram_buffer = addr + total_transferred_bytes_; - int cuda_device = -1; - -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - cuda_device = getCudaDeviceId(addr); - if (cuda_device >= 0) { - dram_buffer = new char[buffer_size]; - cudaSetDevice(cuda_device); -#ifdef USE_MACA - cudaError_t cuda_status = copyTcpCudaMemory( - dram_buffer, addr + total_transferred_bytes_, buffer_size); -#else - cudaError_t cuda_status = - cudaMemcpy(dram_buffer, addr + total_transferred_bytes_, - buffer_size, cudaMemcpyDefault); -#endif - if (cuda_status != cudaSuccess) { - LOG(ERROR) << "ServerSession::writeBody failed to copy from " - "CUDA memory. " - << "Error: " << cudaGetErrorString(cuda_status); - delete[] dram_buffer; - return; // Connection will be closed - } - } -#endif +void tcpTransportSetLaneConnectHandlerHookForTest( + LaneConnectHandlerHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_connect_handler_hook = hook; +} - asio::async_write( - *socket_, asio::buffer(dram_buffer, buffer_size), - [this, addr, dram_buffer, cuda_device, self]( - const asio::error_code& ec, std::size_t transferred_bytes) { -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - if (cuda_device >= 0) { - delete[] dram_buffer; - } -#endif - if (ec) { - LOG(ERROR) - << "ServerSession::writeBody failed. " - << "Attempt to write data " << static_cast(addr) - << " using buffer " << static_cast(dram_buffer) - << ". Error: " << ec.message() - << " (value: " << ec.value() << ")"; - return; // Connection will be closed - } - total_transferred_bytes_ += transferred_bytes; - writeBody(); - }); - } +void tcpTransportSetLaneConnectFailureInjectionHookForTest( + LaneConnectFailureInjectionHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_connect_failure_injection_hook = hook; +} - void readBody() { - auto self(shared_from_this()); - uint64_t size = le64toh(header_.size); - char* addr = local_buffer_; - - size_t buffer_size = - std::min(getChunkSize(), size - total_transferred_bytes_); - if (buffer_size == 0) { - // Destination memory now holds the complete payload. Under v2, - // acknowledge before accepting the next request — this is what - // makes the initiator's COMPLETED mean "applied at the - // destination" rather than "left my socket buffer". - if (v2_) { - sendStatus(kStatusOk, [this] { start(); }); - } else { - start(); - } - return; - } +void tcpTransportSetLaneObserverHookForTest(LaneObserverHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_observer_hook = hook; +} - char* dram_buffer = addr + total_transferred_bytes_; - int cuda_device = -1; +void tcpTransportSetLaneRetryHandlerHookForTest( + LaneRetryHandlerHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_retry_handler_hook = hook; +} -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - cuda_device = getCudaDeviceId(addr); - if (cuda_device >= 0) { - dram_buffer = new char[buffer_size]; - } -#endif +void tcpTransportSetLaneAdmissionHandlerHookForTest( + LaneAdmissionHandlerHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_admission_handler_hook = hook; +} - asio::async_read( - *socket_, asio::buffer(dram_buffer, buffer_size), - [this, addr, dram_buffer, cuda_device, self]( - const asio::error_code& ec, std::size_t transferred_bytes) { - if (ec) { - // If client closed connection (EOF), this is normal - don't - // log - if (ec.value() != asio::error::eof) { - LOG(WARNING) - << "ServerSession::readBody failed. " - << "Attempt to read data " - << static_cast(addr) << " using buffer " - << static_cast(dram_buffer) - << ". Error: " << ec.message() - << " (value: " << ec.value() << ")"; - } - if (cuda_device >= 0) delete[] dram_buffer; - return; // Connection will be closed - } +void tcpTransportSetLaneFailureReasonHookForTest( + LaneFailureReasonHook hook) noexcept { + std::lock_guard lock(lane_test_hook_mutex); + lane_failure_reason_hook = hook; +} -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - if (cuda_device >= 0) { - cudaSetDevice(cuda_device); -#ifdef USE_MACA - cudaError_t cuda_status = - copyTcpCudaMemory(addr + total_transferred_bytes_, - dram_buffer, transferred_bytes); -#else - cudaError_t cuda_status = - cudaMemcpy(addr + total_transferred_bytes_, dram_buffer, - transferred_bytes, cudaMemcpyDefault); -#endif - if (cuda_status != cudaSuccess) { - LOG(ERROR) - << "ServerSession::readBody failed to copy to CUDA " - "memory. " - << "Error: " << cudaGetErrorString(cuda_status); - delete[] dram_buffer; - return; // Connection will be closed - } - delete[] dram_buffer; - } +bool tcpTransportLaneTypesAreMoveOnlyForTest() noexcept { + return std::is_move_constructible::value && + !std::is_copy_constructible::value && + !std::is_copy_assignable::value && + std::is_move_constructible::value && + !std::is_copy_constructible::value && + !std::is_copy_assignable::value; +} #endif - total_transferred_bytes_ += transferred_bytes; - readBody(); - }); - } -}; - -// Client-side session: initiates one transfer request -struct ClientSession : public std::enable_shared_from_this { - explicit ClientSession(std::shared_ptr socket, bool use_v2, - std::function on_complete = nullptr) - : socket_(std::move(socket)), - v2_(use_v2), - on_complete_(std::move(on_complete)) {} - - std::shared_ptr socket_; - SessionHeader header_; - uint64_t total_transferred_bytes_; - char* local_buffer_; - bool v2_; - uint64_t status_frame_; - // v2 WRITE runs the body stream and the ack read concurrently (one - // async op per direction; handlers serialize on the io thread). The - // concurrent read lets a rejection — or a v1 server's bogus payload — - // abort a large in-flight WRITE instead of deadlocking on mutually - // full socket buffers, and delivers rejection frames before the close. - bool write_body_done_ = false; - bool write_acked_ok_ = false; - // An early negative/malformed ack can arrive while asio::async_write still - // owns a buffer pointing into the caller's source memory. Do not publish a - // terminal status until that body operation has completed or been - // cancelled: callers are allowed to release the source buffer as soon as - // the transfer becomes terminal. - bool write_body_in_flight_ = false; - bool write_abort_requested_ = false; - // A v2 status frame is prompt by construction: a READ status precedes - // any payload, and a WRITE ack follows at most one chunk's apply after - // the body is done. The only peer that never sends one is a legacy - // server reached through a stale v2 descriptor — and for requests - // shorter than a frame it also keeps the connection open (it streamed - // size < 8 bytes of "READ payload" and is waiting for our next header), - // so without a deadline both sides wait forever. Bound that wait; the - // default is generous so no healthy slow path can trip it, since it only - // covers the frame itself, never payload streaming. - static int statusFrameTimeoutSec() { - const char* env = std::getenv("MC_TCP_STATUS_TIMEOUT_SEC"); - if (env) { - int v = std::atoi(env); - if (v > 0) return v; - } - return 30; - } - std::optional status_timer_; - bool status_deadline_disarmed_ = false; - std::function on_finalize_; - // Invoked exactly once per request with clean=true iff the protocol - // exchange terminated in a well-defined connection state. A socket whose - // request did not end cleanly must not be reused: the server-side session - // may be mid-frame, and the next request's header would be consumed as - // body bytes. - std::function on_complete_; - - void initiate(void* buffer, uint64_t dest_addr, size_t size, - TransferRequest::OpCode opcode) { - local_buffer_ = (char*)buffer; - header_.addr = htole64(dest_addr); - header_.size = htole64(size); - header_.opcode = (uint8_t)opcode | (v2_ ? kOpcodeV2Flag : 0); - total_transferred_bytes_ = 0; - writeHeader(); - } - private: - // All handlers run on the transport's single io thread, so arm/cancel - // and the expiry handler never race. Expiry only closes the socket: the - // pending status read then completes with an error and its handler owns - // the failure path (including source-buffer quiescence for WRITE). - void armStatusDeadline() { - auto self(shared_from_this()); - status_deadline_disarmed_ = false; - status_timer_.emplace(socket_->get_executor()); - status_timer_->expires_after( - std::chrono::seconds(statusFrameTimeoutSec())); - status_timer_->async_wait([this, self](const asio::error_code& ec) { - // The disarmed flag also covers an expiry that was already - // queued when cancel() ran (cancel cannot revoke those, and by - // then the socket may have been re-pooled). - if (ec == asio::error::operation_aborted || - status_deadline_disarmed_) { - return; - } - LOG(ERROR) << "ClientSession: no status frame within " - << statusFrameTimeoutSec() - << "s (peer likely speaks the legacy protocol); " - "dropping connection"; - if (socket_ && socket_->is_open()) { - asio::error_code cec; - socket_->close(cec); - } - }); - } +#include "tcp_transport_session_impl.h" - void cancelStatusDeadline() { - status_deadline_disarmed_ = true; - if (status_timer_) status_timer_->cancel(); - } +namespace { +constexpr size_t kMaxTcpLanesPerPeer = 16; - // Single terminal path: finish connection ownership, then report the - // outcome. Posted so it runs after the invoking callback returns. - void finalize(TransferStatusEnum status, bool clean) { - cancelStatusDeadline(); - auto self(shared_from_this()); - asio::post( - socket_->get_executor(), - [this, self, status, clean, on_finalize = std::move(on_finalize_), - on_complete = std::move(on_complete_)]() { - // Finish connection ownership before publishing terminal - // status. Once on_finalize marks the slice, the caller may - // immediately free the batch or destroy the transport. - if (on_complete) on_complete(clean); - if (on_finalize) on_finalize(status); - }); - } +size_t parseBoundedTcpSetting(const char* name, const char* value, + size_t default_value, size_t minimum, + size_t maximum) { + if (!value) return default_value; - // Abort a v2 WRITE and cancel any body operation. If asio still owns the - // current source buffer, its completion handler is responsible for - // finalizing after the buffer is quiescent. - void abortWrite() { - write_abort_requested_ = true; - if (socket_ && socket_->is_open()) { - asio::error_code ec; - socket_->close(ec); + const std::string text(value); + size_t parsed = 0; + bool valid = !text.empty(); + for (char c : text) { + if (c < '0' || c > '9') { + valid = false; + break; } - if (!write_body_in_flight_) finalize(TransferStatusEnum::FAILED, false); - } - - void writeHeader() { - auto self(shared_from_this()); - asio::async_write( - *socket_, asio::buffer(&header_, sizeof(SessionHeader)), - [this, self](const asio::error_code& ec, std::size_t len) { - if (ec || len != sizeof(SessionHeader)) { - LOG(ERROR) - << "ClientSession::writeHeader failed. Error: " - << ec.message() << " (value: " << ec.value() << ")" - << ", bytes written: " << len; - finalize(TransferStatusEnum::FAILED, false); - return; - } - if ((header_.opcode & ~kOpcodeV2Flag) == - (uint8_t)TransferRequest::WRITE) { - if (v2_) readWriteAck(); // concurrent with the body - writeBody(); - } else if (v2_) { - readReadStatus(); - } else { - readBody(); - } - }); - } - - // v2 READ: the server prefixes the data with a status frame. - void readReadStatus() { - auto self(shared_from_this()); - armStatusDeadline(); - asio::async_read( - *socket_, asio::buffer(&status_frame_, sizeof(status_frame_)), - [this, self](const asio::error_code& ec, std::size_t len) { - cancelStatusDeadline(); - if (ec || len != sizeof(status_frame_)) { - LOG(ERROR) - << "ClientSession: failed to read READ status " - "frame. Error: " - << ec.message() << " (value: " << ec.value() << ")"; - finalize(TransferStatusEnum::FAILED, false); - return; - } - uint64_t frame = le64toh(status_frame_); - if (!statusFrameValid(frame)) { - LOG(ERROR) << "ClientSession: malformed READ status " - "frame (peer likely speaks the legacy " - "protocol); dropping connection"; - finalize(TransferStatusEnum::FAILED, false); - return; - } - if (frame != kStatusOk) { - LOG(ERROR) << "ClientSession: READ rejected by server, " - "status " - << (frame & 0xFFFFFFFFull); - finalize(TransferStatusEnum::FAILED, false); - return; - } - readBody(); - }); - } - - // v2 WRITE: completion is the server's acknowledgment that the payload - // has been applied to destination memory. Armed concurrently with the - // body stream; a well-behaved v2 server only sends the frame after the - // final chunk, so a frame arriving before the body is done is either a - // rejection or a legacy peer's payload — both close the socket immediately - // and publish failure only after the outstanding body write is quiescent. - void readWriteAck() { - auto self(shared_from_this()); - asio::async_read( - *socket_, asio::buffer(&status_frame_, sizeof(status_frame_)), - [this, self](const asio::error_code& ec, std::size_t len) { - cancelStatusDeadline(); - if (ec || len != sizeof(status_frame_)) { - // The body path may have already finalized a failure and - // closed the socket; finalize() is idempotent (moved-from - // callbacks are null-checked). - if (ec != asio::error::operation_aborted) { - LOG(ERROR) - << "ClientSession: failed to read WRITE " - "ack frame. Error: " - << ec.message() << " (value: " << ec.value() << ")"; - } - abortWrite(); - return; - } - uint64_t frame = le64toh(status_frame_); - if (!statusFrameValid(frame)) { - LOG(ERROR) << "ClientSession: malformed WRITE ack frame " - "(peer likely speaks the legacy protocol); " - "dropping connection"; - abortWrite(); - return; - } - if (frame != kStatusOk) { - LOG(ERROR) << "ClientSession: WRITE rejected by server, " - "status " - << (frame & 0xFFFFFFFFull); - abortWrite(); - return; - } - if (!write_body_done_) { - // The server's ack can legitimately overtake the final - // local write-completion handler (both become ready - // together for small writes; the io thread may run this - // handler first). Record it; the body path finalizes. - write_acked_ok_ = true; - return; - } - finalize(TransferStatusEnum::COMPLETED, true); - }); - } - - void readBody() { - auto self(shared_from_this()); - uint64_t size = le64toh(header_.size); - char* addr = local_buffer_; - - size_t buffer_size = - std::min(getChunkSize(), size - total_transferred_bytes_); - if (buffer_size == 0) { - finalize(TransferStatusEnum::COMPLETED, true); - return; + const size_t digit = static_cast(c - '0'); + if (parsed > (maximum - digit) / size_t(10)) { + valid = false; + break; } - - char* dram_buffer = addr + total_transferred_bytes_; - int cuda_device = -1; - -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - cuda_device = getCudaDeviceId(addr); - if (cuda_device >= 0) { - dram_buffer = new char[buffer_size]; - } -#endif - - asio::async_read( - *socket_, asio::buffer(dram_buffer, buffer_size), - [this, addr, dram_buffer, cuda_device, self]( - const asio::error_code& ec, std::size_t transferred_bytes) { - if (ec) { - LOG(ERROR) - << "ClientSession::readBody failed. " - << "Attempt to read data " << static_cast(addr) - << " using buffer " << static_cast(dram_buffer) - << ". Error: " << ec.message() - << " (value: " << ec.value() << ")"; -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - if (cuda_device >= 0) delete[] dram_buffer; -#endif - finalize(TransferStatusEnum::FAILED, false); - return; - } - -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - if (cuda_device >= 0) { - cudaSetDevice(cuda_device); -#ifdef USE_MACA - cudaError_t cuda_status = - copyTcpCudaMemory(addr + total_transferred_bytes_, - dram_buffer, transferred_bytes); -#else - cudaError_t cuda_status = - cudaMemcpy(addr + total_transferred_bytes_, dram_buffer, - transferred_bytes, cudaMemcpyDefault); -#endif - if (cuda_status != cudaSuccess) { - LOG(ERROR) - << "ClientSession::readBody failed to copy to CUDA " - "memory. " - << "Error: " << cudaGetErrorString(cuda_status); - delete[] dram_buffer; - finalize(TransferStatusEnum::FAILED, false); - return; - } - delete[] dram_buffer; - } -#endif - total_transferred_bytes_ += transferred_bytes; - readBody(); - }); + parsed = parsed * 10 + digit; } + if (valid && parsed >= minimum && parsed <= maximum) return parsed; - void writeBody() { - auto self(shared_from_this()); - uint64_t size = le64toh(header_.size); - char* addr = local_buffer_; - - size_t buffer_size = - std::min(getChunkSize(), size - total_transferred_bytes_); - if (buffer_size == 0) { - if (v2_) { - if (write_abort_requested_) { - finalize(TransferStatusEnum::FAILED, false); - return; - } - // Completion comes from the server's acknowledgment, whose - // read is already in flight (armed in writeHeader) and may - // have finished first. - write_body_done_ = true; - if (write_acked_ok_) { - finalize(TransferStatusEnum::COMPLETED, true); - } else { - // From here a well-behaved server owes at most one - // chunk's apply plus the frame; a legacy peer behind a - // stale descriptor may owe nothing, ever. - armStatusDeadline(); - } - } else { - // v1: no acknowledgment exists in the protocol; this only - // means the payload left the initiator (#2086). - finalize(TransferStatusEnum::COMPLETED, true); - } - return; - } - - char* dram_buffer = addr + total_transferred_bytes_; - int cuda_device = -1; - -#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ - defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ - defined(USE_COREX) - cuda_device = getCudaDeviceId(addr); - if (cuda_device >= 0) { - dram_buffer = new char[buffer_size]; - cudaSetDevice(cuda_device); -#ifdef USE_MACA - cudaError_t cuda_status = copyTcpCudaMemory( - dram_buffer, addr + total_transferred_bytes_, buffer_size); -#else - cudaError_t cuda_status = - cudaMemcpy(dram_buffer, addr + total_transferred_bytes_, - buffer_size, cudaMemcpyDefault); -#endif - if (cuda_status != cudaSuccess) { - LOG(ERROR) << "ClientSession::writeBody failed to copy from " - "CUDA memory. " - << "Error: " << cudaGetErrorString(cuda_status); - delete[] dram_buffer; - abortWrite(); - return; - } - } -#endif - - write_body_in_flight_ = true; - asio::async_write( - *socket_, asio::buffer(dram_buffer, buffer_size), - [this, addr, dram_buffer, cuda_device, self]( - const asio::error_code& ec, std::size_t transferred_bytes) { - write_body_in_flight_ = false; - if (cuda_device >= 0) { - delete[] dram_buffer; - } - if (ec) { - LOG(ERROR) - << "ClientSession::writeBody failed. " - << "Attempt to write data " << static_cast(addr) - << " using buffer " << static_cast(dram_buffer) - << ". Error: " << ec.message() - << " (value: " << ec.value() << ")"; - abortWrite(); - return; - } - if (write_abort_requested_) { - // The early ack path closed the socket while this - // operation still owned the caller's source buffer. It is - // safe to publish failure now that the handler has run. - finalize(TransferStatusEnum::FAILED, false); - return; - } - total_transferred_bytes_ += transferred_bytes; - writeBody(); - }); - } -}; + LOG(WARNING) << "Invalid " << name << " value: " << text + << ", using default " << default_value; + return default_value; +} -struct TcpContext { - TcpContext(short port, ValidateAddrFn validate_addr) - : acceptor(io_context), validate_addr_(std::move(validate_addr)) { - std::error_code ec; - asio::ip::tcp::endpoint endpoint(asio::ip::tcp::v6(), port); - - acceptor.open(endpoint.protocol(), ec); - if (!ec) { - acceptor.set_option(asio::ip::v6_only(false), ec); - if (!ec) { - acceptor.set_option( - asio::ip::tcp::acceptor::reuse_address(true)); - acceptor.bind(endpoint, ec); - if (!ec) { - acceptor.listen(); - return; - } - } - acceptor.close(); - } - LOG(ERROR) << "Failed to set up IPv6 dual-stack listener: " - << ec.message() << " (error code: " << ec.value() << ")"; - asio::ip::tcp::endpoint endpoint_v4(asio::ip::tcp::v4(), port); - acceptor.open(endpoint_v4.protocol()); - acceptor.set_option(asio::ip::tcp::acceptor::reuse_address(true)); - acceptor.bind(endpoint_v4); - acceptor.listen(); - } +bool validateTcpAddress(const std::shared_ptr& metadata, + uint64_t addr, uint64_t size) { + if (size == 0 || addr + size < addr) return false; - void doAccept() { - acceptor.async_accept([this](asio::error_code ec, tcpsocket socket) { - if (!ec) { - asio::error_code nodelay_ec; - socket.set_option(asio::ip::tcp::no_delay(true), nodelay_ec); - auto socket_ptr = - std::make_shared(std::move(socket)); - auto session = - std::make_shared(socket_ptr, validate_addr_); - session->start(); - } - doAccept(); - }); + auto desc = metadata->getSegmentDescByID(LOCAL_SEGMENT_ID); + if (!desc) return false; + for (const auto& buffer : desc->buffers) { + if (buffer.addr + buffer.length < buffer.addr) continue; + if (buffer.addr <= addr && addr + size <= buffer.addr + buffer.length) + return true; } + return false; +} +} // namespace - asio::io_context io_context; - asio::ip::tcp::acceptor acceptor; - ValidateAddrFn validate_addr_; -}; - -TcpTransport::TcpTransport() : context_(nullptr), running_(false) { +TcpTransport::TcpTransport() + : context_(nullptr), + running_(false), + lane_state_(std::make_shared()) { if (getenv("MC_TCP_ENABLE_CONNECTION_POOL") != nullptr) { std::string val(getenv("MC_TCP_ENABLE_CONNECTION_POOL")); std::transform(val.begin(), val.end(), val.begin(), @@ -851,22 +250,49 @@ TcpTransport::TcpTransport() : context_(nullptr), running_(false) { enable_connection_pool_ = true; } } -} -TcpTransport::~TcpTransport() { - if (running_) { - running_ = false; - context_->io_context.stop(); - thread_.join(); - } + constexpr size_t kDefaultLanesPerPeer = 4; + constexpr size_t kDefaultQueuedTransfersPerPeer = 1024; + constexpr size_t kMaxQueuedTransfersPerPeer = 65535; + constexpr size_t kDefaultPendingAdmissionsPerPeer = 1024; + constexpr size_t kMaxPendingAdmissionsPerPeer = 65535; + constexpr size_t kDefaultAdmissionTimeoutMs = 1000; + constexpr size_t kMaxAdmissionTimeoutMs = 600000; - // Clear connection pool BEFORE deleting context - // because sockets in the pool reference io_context - { - std::lock_guard lock(pool_mutex_); - connection_pool_.clear(); + const char* lanes_env = getenv("MC_TCP_LANES_PER_PEER"); + if (lanes_env) { + lane_state_->lanes_per_peer = parseBoundedTcpSetting( + "MC_TCP_LANES_PER_PEER", lanes_env, kDefaultLanesPerPeer, 1, + kMaxTcpLanesPerPeer); + } else { + const char* deprecated_env = getenv("MC_TCP_MAX_CONNECTIONS_PER_PEER"); + if (deprecated_env) { + LOG(WARNING) << "MC_TCP_MAX_CONNECTIONS_PER_PEER is deprecated; " + "use MC_TCP_LANES_PER_PEER"; + lane_state_->lanes_per_peer = parseBoundedTcpSetting( + "MC_TCP_MAX_CONNECTIONS_PER_PEER", deprecated_env, + kDefaultLanesPerPeer, 1, kMaxTcpLanesPerPeer); + } } + lane_state_->max_queued_transfers_per_peer = parseBoundedTcpSetting( + "MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", + getenv("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER"), + kDefaultQueuedTransfersPerPeer, 1, kMaxQueuedTransfersPerPeer); + lane_state_->max_pending_admissions_per_peer = parseBoundedTcpSetting( + "MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + getenv("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER"), + kDefaultPendingAdmissionsPerPeer, 1, kMaxPendingAdmissionsPerPeer); + lane_state_->admission_timeout = + std::chrono::milliseconds(parseBoundedTcpSetting( + "MC_TCP_ADMISSION_TIMEOUT_MS", + getenv("MC_TCP_ADMISSION_TIMEOUT_MS"), kDefaultAdmissionTimeoutMs, + 1, kMaxAdmissionTimeoutMs)); +} + +TcpTransport::~TcpTransport() { + shutdownConnectionLanes(); + if (context_) { delete context_; context_ = nullptr; @@ -915,9 +341,14 @@ int TcpTransport::install(std::string& local_server_name, close(sockfd); // the above function has opened a socket LOG(INFO) << "TcpTransport: listen on port " << tcp_port; - context_ = new TcpContext(tcp_port, [this](uint64_t addr, uint64_t size) { - return validateAddress(addr, size); + auto metadata = metadata_; + context_ = new TcpContext(tcp_port, [metadata = std::move(metadata)]( + uint64_t addr, uint64_t size) { + return validateTcpAddress(metadata, addr, size); }); + lane_runtime_ = + std::make_shared(context_->io_context); + lane_state_->runtime = lane_runtime_; running_ = true; thread_ = std::thread(&TcpTransport::worker, this); return 0; @@ -1072,22 +503,69 @@ Transport::Slice* TcpTransport::prepareTransfer( void TcpTransport::startTransferSequence(std::vector slices) { struct Sequence { + std::mutex mutex; std::vector slices; size_t next = 0; + bool advancing = false; + bool resume_requested = false; }; + auto sequence = std::make_shared(); sequence->slices = std::move(slices); + auto advance = std::make_shared>(); std::weak_ptr> weak_advance = advance; + *advance = [this, sequence, weak_advance]() { auto advance = weak_advance.lock(); - if (!advance || sequence->next == sequence->slices.size()) return; - auto* slice = sequence->slices[sequence->next++]; - std::function continuation; - if (sequence->next < sequence->slices.size()) - continuation = [advance]() { (*advance)(); }; - startTransfer(slice, std::move(continuation), true); + if (!advance) return; + + { + std::lock_guard lock(sequence->mutex); + if (sequence->next == sequence->slices.size()) return; + if (sequence->advancing) { + sequence->resume_requested = true; + return; + } + sequence->advancing = true; + } + + while (true) { + Slice* slice = nullptr; + bool has_more = false; + { + std::lock_guard lock(sequence->mutex); + if (sequence->next == sequence->slices.size()) { + sequence->advancing = false; + return; + } + slice = sequence->slices[sequence->next++]; + has_more = sequence->next < sequence->slices.size(); + sequence->resume_requested = false; + } + + std::function continuation; + if (has_more) { + continuation = [advance]() { (*advance)(); }; + } + + // startTransfer may fail synchronously (including during shutdown). + // The trampoline above turns a synchronous continuation into + // another loop iteration rather than recursive calls. For an + // asynchronous completion, this invocation returns and the + // continuation becomes the next runner. + startTransfer(slice, std::move(continuation), true); + + { + std::lock_guard lock(sequence->mutex); + if (!sequence->resume_requested) { + sequence->advancing = false; + return; + } + } + } }; + (*advance)(); } @@ -1106,220 +584,41 @@ void TcpTransport::worker() { } std::shared_ptr TcpTransport::getConnection( - const std::string& host, uint16_t port, bool use_pool) { - // Ungrouped transfers keep the configured connection-pool behavior. - if (!use_pool) { - try { - asio::ip::tcp::resolver resolver(context_->io_context); - auto endpoint_iterator = - resolver.resolve(host, std::to_string(port)); - auto socket_ptr = - std::make_shared(context_->io_context); - asio::connect(*socket_ptr, endpoint_iterator); - socket_ptr->set_option(asio::ip::tcp::no_delay(true)); - return socket_ptr; - } catch (std::exception& e) { - LOG(ERROR) - << "TcpTransport::getConnection failed to create connection to " - << host << ":" << port << ". Error: " << e.what(); - return nullptr; - } - } - - ConnectionKey key{host, port}; - - // First phase: search for available connection while holding the lock - { - std::lock_guard lock(pool_mutex_); - - // Cleanup idle and dead connections - cleanupIdleConnections(); - - auto it = connection_pool_.find(key); - if (it != connection_pool_.end()) { - auto& queue = it->second; - - // Find an available connection - for (auto queue_it = queue.begin(); queue_it != queue.end();) { - auto& entry = *queue_it; - if (!entry->in_use) { - // Check if connection is still alive - if (entry->socket->is_open()) { - entry->in_use = true; - entry->last_used = std::chrono::steady_clock::now(); - return entry->socket; - } else { - // Remove dead connection immediately - queue_it = queue.erase(queue_it); - continue; - } - } - ++queue_it; - } - } - } - - // No available connection, create a new one (pool grows dynamically) - // Release lock before creating new connection to avoid blocking other - // threads during slow DNS resolution and TCP handshake - std::shared_ptr new_socket; + const std::string& host, uint16_t port) { + // The reusable path is owned by fixed connection lanes. This helper is + // only for the connection-pool-disabled one-shot path. try { asio::ip::tcp::resolver resolver(context_->io_context); auto endpoint_iterator = resolver.resolve(host, std::to_string(port)); - new_socket = + auto socket_ptr = std::make_shared(context_->io_context); - asio::connect(*new_socket, endpoint_iterator); - new_socket->set_option(asio::ip::tcp::no_delay(true)); + asio::connect(*socket_ptr, endpoint_iterator); + socket_ptr->set_option(asio::ip::tcp::no_delay(true)); + return socket_ptr; } catch (std::exception& e) { LOG(ERROR) << "TcpTransport::getConnection failed to create connection to " << host << ":" << port << ". Error: " << e.what(); return nullptr; } - - // Re-acquire lock to add the new connection to the pool - std::shared_ptr entry; - { - std::lock_guard lock(pool_mutex_); - // Re-check if another thread already added a connection while we were - // creating this one - auto& queue = connection_pool_[key]; - for (auto it = queue.begin(); it != queue.end(); ++it) { - auto& existing_entry = *it; - if (!existing_entry->in_use && existing_entry->socket->is_open()) { - // Another thread added an available connection, use that - // instead and close the one we just created - if (new_socket && new_socket->is_open()) { - asio::error_code ec; - new_socket->close(ec); - } - existing_entry->in_use = true; - existing_entry->last_used = std::chrono::steady_clock::now(); - return existing_entry->socket; - } - } - - // No other connection available, add the one we created to the pool - entry = std::make_shared(new_socket, host, port); - queue.push_back(entry); - } - - return entry->socket; -} - -void TcpTransport::returnConnection( - const std::string& host, uint16_t port, - std::shared_ptr socket) { - ConnectionKey key{host, port}; - - std::lock_guard lock(pool_mutex_); - - auto it = connection_pool_.find(key); - if (it != connection_pool_.end()) { - for (auto entry_it = it->second.begin(); entry_it != it->second.end(); - ++entry_it) { - if ((*entry_it)->socket == socket) { - if (socket->is_open()) { - (*entry_it)->in_use = false; - (*entry_it)->last_used = std::chrono::steady_clock::now(); - } else { - // Connection is dead, remove from pool - it->second.erase(entry_it); - } - return; - } - } - } - - // Connection not found in pool (might be temporary), close it - if (socket && socket->is_open()) { - asio::error_code ec; - socket->close(ec); - } -} - -void TcpTransport::cleanupIdleConnections() { - auto now = std::chrono::steady_clock::now(); - - for (auto it = connection_pool_.begin(); it != connection_pool_.end();) { - auto& queue = it->second; - - for (auto entry_it = queue.begin(); entry_it != queue.end();) { - auto& entry = *entry_it; - if (!entry->in_use) { - auto idle_duration = - std::chrono::duration_cast( - now - entry->last_used) - .count(); - if (idle_duration > kConnectionIdleTimeout.count()) { - if (entry->socket && entry->socket->is_open()) { - asio::error_code ec; - entry->socket->close(ec); - } - entry_it = queue.erase(entry_it); - continue; - } - } - ++entry_it; - } - - if (queue.empty()) { - it = connection_pool_.erase(it); - } else { - ++it; - } - } } -bool TcpTransport::validateAddress(uint64_t addr, uint64_t size) const { - if (size == 0) return false; - if (addr + size < addr) return false; - - auto desc = metadata_->getSegmentDescByID(LOCAL_SEGMENT_ID); - if (!desc) return false; - - for (const auto& buffer : desc->buffers) { - if (buffer.addr + buffer.length < buffer.addr) continue; - if (buffer.addr <= addr && addr + size <= buffer.addr + buffer.length) - return true; - } - return false; -} - -void TcpTransport::discardConnection( - const std::string& host, uint16_t port, - std::shared_ptr socket) { - if (socket && socket->is_open()) { - asio::error_code ec; - socket->close(ec); - } - std::lock_guard lock(pool_mutex_); - auto it = connection_pool_.find(ConnectionKey{host, port}); - if (it != connection_pool_.end()) { - auto& queue = it->second; - for (auto queue_it = queue.begin(); queue_it != queue.end(); - ++queue_it) { - if ((*queue_it)->socket == socket) { - queue.erase(queue_it); - break; - } - } - if (queue.empty()) connection_pool_.erase(it); - } -} +#include "tcp_transport_lane_impl.h" void TcpTransport::startTransfer(Slice* slice, std::function continuation, bool reuse_connection) { - auto finish = [this, slice, continuation = std::move(continuation)]( - TransferStatusEnum status) mutable { + auto finish = [slice, &continuation](TransferStatusEnum status) mutable { if (status == TransferStatusEnum::COMPLETED) slice->markSuccess(); else slice->markFailed(); - if (continuation) - asio::post(context_->io_context, std::move(continuation)); + if (continuation) { + auto next = std::move(continuation); + next(); + } }; + auto desc = metadata_->getSegmentDescByID(slice->target_id); if (!desc) { LOG(ERROR) << "TcpTransport::startTransfer failed to get segment " @@ -1340,68 +639,91 @@ void TcpTransport::startTransfer(Slice* slice, // Zero-length requests are complete by definition. v1 reported them // COMPLETED while the server silently rejected size==0 in address - // validation; short-circuiting keeps that outcome (rather than turning - // no-ops into v2 rejection failures) without the pointless round trip. + // validation; preserve that outcome without a round trip. if (slice->length == 0) { finish(TransferStatusEnum::COMPLETED); return; } - const bool use_pool = enable_connection_pool_ || reuse_connection; - auto socket = getConnection(meta_entry.ip_or_host_name, desc->tcp_data_port, - use_pool); + const ConnectionKey key{meta_entry.ip_or_host_name, + static_cast(desc->tcp_data_port)}; + const bool use_v2 = desc->tcp_proto_version >= 2 && !forceLegacyTcpProto(); + TcpWorkItem work(slice, use_v2, std::move(continuation)); + + // Scatter task groups request reuse even when the general pool setting is + // disabled. Fixed lanes provide the same serial socket reuse without + // reviving the old unbounded dynamic pool. + if (enable_connection_pool_ || reuse_connection) { + enqueuePooledTransfer(key, std::move(work)); + return; + } + + // Preserve the connection-pool-disabled synchronous one-shot path. + auto socket = getConnection(key.host, key.port); if (!socket) { LOG(ERROR) << "TcpTransport::startTransfer failed to get connection to " - << meta_entry.ip_or_host_name << ":" << desc->tcp_data_port; - finish(TransferStatusEnum::FAILED); + << key.host << ":" << key.port; + completeTerminalAction( + TerminalAction(std::move(work), TransferStatusEnum::FAILED, false)); return; } + startTransferWithSocket(std::move(work), std::move(socket)); +} +void TcpTransport::startTransferWithSocket( + TcpWorkItem work, std::shared_ptr socket) noexcept { + const Slice* slice = work.slice; + std::shared_ptr> terminal_work; try { - const bool use_v2 = - desc->tcp_proto_version >= 2 && !forceLegacyTcpProto(); - auto session = std::make_shared(socket, use_v2); - - session->on_finalize_ = finish; - - // Return connection to pool when the request terminated cleanly; - // otherwise the server-side session state is unknown (it may be - // mid-frame), so reusing the socket would desynchronize the next - // request. Discard it instead. - if (use_pool) { - session->on_complete_ = [this, host = meta_entry.ip_or_host_name, - port = desc->tcp_data_port, - socket](bool clean) { - if (clean) - returnConnection(host, port, socket); - else - discardConnection(host, port, socket); - }; - } else { - session->on_complete_ = [socket](bool) { - // Close connection immediately after transfer - if (socket && socket->is_open()) { - asio::error_code ec; - socket->close(ec); - } - }; - } - - session->initiate(slice->source_addr, slice->tcp.dest_addr, - slice->length, slice->opcode); - } catch (std::exception& e) { + terminal_work = + std::make_shared>(std::move(work)); + auto session = std::make_shared( + socket, terminal_work->value().use_v2, + [terminal_work, socket](TransferStatusEnum status, bool) noexcept { + closeSocketNoThrow(socket); + if (!terminal_work->has_value()) return; + auto completed = std::move(terminal_work->value()); + terminal_work->reset(); + completeTerminalAction( + TerminalAction(std::move(completed), status, false)); + }); + session->initiate(terminal_work->value().slice->source_addr, + terminal_work->value().slice->tcp.dest_addr, + terminal_work->value().slice->length, + terminal_work->value().slice->opcode); + } catch (const std::exception& e) { LOG(ERROR) << "TcpTransport::startTransfer encountered an exception. " "Slice details - source_addr: " << slice->source_addr << ", length: " << slice->length << ", opcode: " << (int)slice->opcode << ", target_id: " << slice->target_id << ". Exception: " << e.what(); - // On exception, always close the socket and remove from pool if - // present. Don't return it to the pool as it may be in an - // inconsistent state. - discardConnection(meta_entry.ip_or_host_name, - static_cast(desc->tcp_data_port), socket); - finish(TransferStatusEnum::FAILED); + closeSocketNoThrow(socket); + if (terminal_work && terminal_work->has_value()) { + auto failed = std::move(terminal_work->value()); + terminal_work->reset(); + failWorkItem(std::move(failed), WorkFailureReason::SESSION_FAILED, + lane_state_->failure_counters); + } else if (work.slice) { + failWorkItem(std::move(work), WorkFailureReason::SESSION_FAILED, + lane_state_->failure_counters); + } + } catch (...) { + LOG(ERROR) << "TcpTransport::startTransfer encountered an unknown " + "exception. Slice details - source_addr: " + << slice->source_addr << ", length: " << slice->length + << ", opcode: " << (int)slice->opcode + << ", target_id: " << slice->target_id; + closeSocketNoThrow(socket); + if (terminal_work && terminal_work->has_value()) { + auto failed = std::move(terminal_work->value()); + terminal_work->reset(); + failWorkItem(std::move(failed), WorkFailureReason::SESSION_FAILED, + lane_state_->failure_counters); + } else if (work.slice) { + failWorkItem(std::move(work), WorkFailureReason::SESSION_FAILED, + lane_state_->failure_counters); + } } } diff --git a/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_lane_impl.h b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_lane_impl.h new file mode 100644 index 0000000000..1180a2fe15 --- /dev/null +++ b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_lane_impl.h @@ -0,0 +1,1390 @@ +// Copyright 2024 KVCache.AI +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Private implementation fragment included by tcp_transport.cpp from inside +// namespace mooncake. This file contains the bounded connection-lane state +// machine and its test hooks; it is intentionally not a standalone translation +// unit so this review-only split does not change linkage or initialization. + +namespace { +constexpr size_t kMaxConcurrentLaneProbes = 1; +// Conservative fixed policy for the first rate-limited implementation. The +// exact cooldown/backoff policy remains subject to maintainer review. +constexpr auto kReconnectRoundCooldown = std::chrono::seconds(1); +constexpr auto kShutdownCancellationWait = std::chrono::seconds(2); + +bool shouldLogOccurrence(uint64_t occurrence) { + return occurrence != 0 && (occurrence & (occurrence - 1)) == 0; +} + +void cancelTimerNoThrow( + const std::shared_ptr& timer) noexcept { + if (!timer) return; + asio::error_code ec; + timer->cancel(ec); +} + +// Tracks only whether executor-posted cancellation actions ran. It is not an +// asynchronous-handler quiescence barrier; stop/join provides that boundary. +struct LaneCancellationPostTracker { + std::mutex mutex; + std::condition_variable cv; + size_t pending = 0; + + void add() { + std::lock_guard lock(mutex); + ++pending; + } + + void done() { + { + std::lock_guard lock(mutex); + if (pending != 0) --pending; + } + cv.notify_all(); + } + + void waitUntil(std::chrono::steady_clock::time_point deadline) { + std::unique_lock lock(mutex); + cv.wait_until(lock, deadline, [this] { return pending == 0; }); + } +}; +} // namespace + +bool TcpTransport::hasUsableLaneLocked(const PeerConnectionGroup& group) { + for (const auto& lane : group.lanes) { + if ((lane->state == LaneState::IDLE || lane->state == LaneState::BUSY || + lane->state == LaneState::COMPLETING) && + lane->socket && lane->socket->is_open()) { + return true; + } + } + return false; +} + +bool TcpTransport::hasDisconnectedLaneLocked(const PeerConnectionGroup& group) { + return std::any_of(group.lanes.begin(), group.lanes.end(), + [](const auto& lane) { + return lane->state == LaneState::DISCONNECTED; + }); +} + +bool TcpTransport::hasUntriedDisconnectedLaneLocked( + const PeerConnectionGroup& group) { + return std::any_of( + group.lanes.begin(), group.lanes.end(), [&group](const auto& lane) { + return lane->state == LaneState::DISCONNECTED && + lane->last_connect_round != group.connect_round; + }); +} + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +size_t TcpTransport::activeSocketCountLocked(const PeerConnectionGroup& group) { + size_t count = 0; + for (const auto& lane : group.lanes) { + if (lane->resolver || lane->socket) ++count; + } + return count; +} +#endif + +void TcpTransport::beginConnectRoundLocked(PeerConnectionGroup& group) { + ++group.connect_round; + if (group.connect_round == 0) group.connect_round = 1; + group.connect_round_had_success = false; + group.next_probe_not_before = {}; +} + +void TcpTransport::enterReconnectCooldownLocked(PeerConnectionGroup& group) { + group.next_probe_not_before = + std::chrono::steady_clock::now() + kReconnectRoundCooldown; +} + +void TcpTransport::addQueuedBytesLocked(PeerConnectionGroup& group, + uint64_t length) { + if (group.queued_bytes_saturated) return; + if (group.queued_bytes > std::numeric_limits::max() - length) { + group.queued_bytes = std::numeric_limits::max(); + group.queued_bytes_saturated = true; + return; + } + group.queued_bytes += length; +} + +void TcpTransport::removeQueuedBytesLocked(PeerConnectionGroup& group, + uint64_t length) { + if (!group.queued_bytes_saturated) { + group.queued_bytes = + group.queued_bytes >= length ? group.queued_bytes - length : 0; + return; + } + + // Once saturated, UINT64_MAX is only an explicit lower-fidelity marker, + // not an exact sum. Recompute after removal so exact accounting resumes as + // soon as the remaining queue fits in uint64_t. + group.queued_bytes = 0; + group.queued_bytes_saturated = false; + for (const auto& item : group.queue) { + addQueuedBytesLocked(group, item.slice->length); + if (group.queued_bytes_saturated) break; + } +} + +void TcpTransport::clearQueuedBytesLocked(PeerConnectionGroup& group) { + group.queued_bytes = 0; + group.queued_bytes_saturated = false; +} + +size_t TcpTransport::expirePendingAdmissionsLocked( + PeerConnectionGroup& group, std::chrono::steady_clock::time_point now, + std::deque& expired) { + size_t count = 0; + while (!group.pending_admissions.empty() && + group.pending_admissions.front().admission_deadline <= now) { + expired.emplace_back(std::move(group.pending_admissions.front())); + group.pending_admissions.pop_front(); + ++count; + } + return count; +} + +size_t TcpTransport::promotePendingAdmissionsLocked( + PeerConnectionGroup& group) { + size_t count = 0; + while (group.queue.size() < group.queue_capacity && + !group.pending_admissions.empty()) { + group.queue.emplace_back(std::move(group.pending_admissions.front())); + group.pending_admissions.pop_front(); + addQueuedBytesLocked(group, group.queue.back().slice->length); + ++count; + } + return count; +} + +void TcpTransport::refreshAdmissionTimerLocked( + const std::shared_ptr& group, + std::deque& runtime_failed, + std::shared_ptr& timer_to_cancel, bool& timer_armed) { + timer_armed = false; + if (group->pending_admissions.empty()) { + if (group->admission_timer) { + ++group->admission_epoch; + if (group->admission_epoch == 0) ++group->admission_epoch; + timer_to_cancel = std::move(group->admission_timer); + } + return; + } + + if (group->admission_timer) return; + + try { + auto timer = std::make_shared(group->executor); + timer->expires_at(group->pending_admissions.front().admission_deadline); + ++group->admission_epoch; + if (group->admission_epoch == 0) ++group->admission_epoch; + const uint64_t admission_epoch = group->admission_epoch; + group->admission_timer = timer; + timer->async_wait([group, timer, admission_epoch](asio::error_code ec) { +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneAdmissionHandlerHook(); +#endif + handleAdmissionTimer(group, timer, admission_epoch, ec); + }); + timer_armed = true; + return; + } catch (...) { + ++group->admission_epoch; + if (group->admission_epoch == 0) ++group->admission_epoch; + if (group->admission_timer) + timer_to_cancel = std::move(group->admission_timer); + while (!group->pending_admissions.empty()) { + runtime_failed.emplace_back( + std::move(group->pending_admissions.front())); + group->pending_admissions.pop_front(); + } + return; + } +} + +void TcpTransport::handleAdmissionTimer( + const std::shared_ptr& group, + const std::shared_ptr& timer, uint64_t admission_epoch, + asio::error_code ec) { + std::deque expired; + std::deque runtime_failed; + std::shared_ptr timer_to_cancel; + uint64_t pump_epoch = 0; + size_t promoted = 0; + [[maybe_unused]] size_t pending_depth = 0; + bool timer_armed = false; + [[maybe_unused]] bool fired = false; + [[maybe_unused]] bool late = false; + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + group->admission_epoch != admission_epoch || + group->admission_timer != timer) { + late = true; + } else { + group->admission_timer.reset(); + fired = true; + if (ec) { + // Cancellation normally increments admission_epoch first and + // is therefore stale. Keep this guard for other Asio timer + // errors delivered while the matching timer is still active. + while (!group->pending_admissions.empty()) { + runtime_failed.emplace_back( + std::move(group->pending_admissions.front())); + group->pending_admissions.pop_front(); + } + } else { + expirePendingAdmissionsLocked( + *group, std::chrono::steady_clock::now(), expired); + promoted = promotePendingAdmissionsLocked(*group); + refreshAdmissionTimerLocked(group, runtime_failed, + timer_to_cancel, timer_armed); + if (promoted != 0) pump_epoch = requestGroupPumpLocked(*group); + } + pending_depth = group->pending_admissions.size(); + } + } + + cancelTimerNoThrow(timer_to_cancel); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (late) { + invokeLaneObserverHook(kLaneAdmissionTimerLate, 0, 0, 0, false); + } else if (fired) { + invokeLaneObserverHook(kLaneAdmissionTimerFired, pending_depth, 0, 0, + false); + } + if (promoted != 0) + invokeLaneObserverHook(kLaneAdmissionPromoted, pending_depth, promoted, + 0, false); + if (timer_armed) + invokeLaneObserverHook(kLaneAdmissionTimerArmed, pending_depth, 0, 0, + false); +#endif + failWorkItems(std::move(expired), WorkFailureReason::QUEUE_TIMEOUT, + group->failure_counters); + failWorkItems(std::move(runtime_failed), + WorkFailureReason::RUNTIME_UNAVAILABLE, + group->failure_counters); + if (pump_epoch != 0) postGroupPump(group, pump_epoch); +} + +bool TcpTransport::armRetryTimerLocked( + const std::shared_ptr& group) { + if (group->retry_timer || group->state != GroupState::OPEN || + group->queue.empty()) { + return true; + } + + try { + auto timer = std::make_shared(group->executor); + timer->expires_at(group->next_probe_not_before); + ++group->retry_epoch; + if (group->retry_epoch == 0) ++group->retry_epoch; + const uint64_t retry_epoch = group->retry_epoch; + group->retry_timer = timer; + // A handler abandoned by io_context.stop() is harmless: it holds only + // group/timer shared state and can only request a new connection round. + timer->async_wait([group, timer, retry_epoch](asio::error_code ec) { +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneRetryHandlerHook(); +#endif + handleRetryTimer(group, timer, retry_epoch, ec); + }); + return true; + } catch (...) { + group->retry_timer.reset(); + ++group->retry_epoch; + if (group->retry_epoch == 0) ++group->retry_epoch; + return false; + } +} + +void TcpTransport::handleRetryTimer( + const std::shared_ptr& group, + const std::shared_ptr& timer, uint64_t retry_epoch, + asio::error_code ec) { + uint64_t pump_epoch = 0; + [[maybe_unused]] bool fired = false; + [[maybe_unused]] bool late = false; + { + std::lock_guard lock(group->mutex); + if (group->retry_epoch != retry_epoch || group->retry_timer != timer) { + late = true; + } else { + group->retry_timer.reset(); + if (group->state != GroupState::OPEN) { + // Shutdown also bumps retry_epoch, so this is normally caught + // by the identity check above. Retain the state guard for a + // handler that observes the transition at this boundary. + late = group->state != GroupState::OPEN; + } else if (!group->queue.empty() && !ec && + group->probes_in_flight == 0 && + hasDisconnectedLaneLocked(*group) && + !hasUntriedDisconnectedLaneLocked(*group)) { + beginConnectRoundLocked(*group); + pump_epoch = requestGroupPumpLocked(*group); + fired = true; + } else if (!group->queue.empty() && group->probes_in_flight == 0) { + pump_epoch = requestGroupPumpLocked(*group); + } + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (late) + invokeLaneObserverHook(kLaneRetryLate, 0, 0, 0, false); + else if (fired) + invokeLaneObserverHook(kLaneRetryFired, 0, 0, 0, false); +#endif + if (pump_epoch != 0) postGroupPump(group, pump_epoch); +} + +uint64_t TcpTransport::requestGroupPumpLocked(PeerConnectionGroup& group) { + if (group.state != GroupState::OPEN || group.pump_scheduled || + group.queue.empty()) { + return 0; + } + group.pump_scheduled = true; + ++group.pump_epoch; + if (group.pump_epoch == 0) ++group.pump_epoch; + return group.pump_epoch; +} + +void TcpTransport::enqueuePooledTransfer(const ConnectionKey& key, + TcpWorkItem work) { + const auto state = lane_state_; + std::shared_ptr group; + std::optional rejected; + std::deque expired; + std::deque runtime_failed; + std::shared_ptr timer_to_cancel; + WorkFailureReason rejection_reason = WorkFailureReason::QUEUE_FULL; + uint64_t pump_epoch = 0; + [[maybe_unused]] size_t promoted = 0; + [[maybe_unused]] size_t pending_depth = 0; + [[maybe_unused]] bool direct_admission = false; + [[maybe_unused]] bool pending_admission = false; + [[maybe_unused]] bool hard_rejection = false; + bool timer_armed = false; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + size_t queue_depth = 0; + uint64_t queued_bytes = 0; + size_t active_sockets = 0; +#endif + + try { + std::lock_guard state_lock(state->mutex); + if (state->shutting_down) { + rejected.emplace(std::move(work)); + rejection_reason = WorkFailureReason::SHUTDOWN; + } else { + auto runtime = state->runtime.lock(); + if (!runtime) { + rejected.emplace(std::move(work)); + rejection_reason = WorkFailureReason::RUNTIME_UNAVAILABLE; + } else { + auto group_it = state->groups.find(key); + if (group_it == state->groups.end()) { + group = std::make_shared( + key, runtime->executor, + state->max_queued_transfers_per_peer, + state->max_pending_admissions_per_peer, + state->admission_timeout, state->failure_counters); + group->lanes.reserve(state->lanes_per_peer); + for (size_t i = 0; i < state->lanes_per_peer; ++i) { + group->lanes.push_back( + std::make_shared(i, group)); + } + auto [inserted_it, inserted] = + state->groups.emplace(key, group); + if (!inserted) group = inserted_it->second; + } else { + group = group_it->second; + } + + std::lock_guard group_lock(group->mutex); + if (group->state != GroupState::OPEN) { + rejected.emplace(std::move(work)); + rejection_reason = WorkFailureReason::SHUTDOWN; + } else { + // Submissions arriving during cooldown remain subject to + // the bounded queue and wait until the retry timer expires + // and a new round begins, so their added latency is bounded + // by the cooldown length. + const auto admission_time = + std::chrono::steady_clock::now(); + expirePendingAdmissionsLocked(*group, admission_time, + expired); + promoted = promotePendingAdmissionsLocked(*group); + + if (group->pending_admissions.empty() && + group->queue.size() < group->queue_capacity) { + group->queue.emplace_back(std::move(work)); + addQueuedBytesLocked(*group, + group->queue.back().slice->length); + direct_admission = true; + } else if (group->pending_admissions.size() < + group->pending_admission_capacity) { + work.admission_deadline = + admission_time + group->admission_timeout; + group->pending_admissions.emplace_back(std::move(work)); + pending_admission = true; + } else { + rejected.emplace(std::move(work)); + hard_rejection = true; + } + + refreshAdmissionTimerLocked(group, runtime_failed, + timer_to_cancel, timer_armed); + pump_epoch = requestGroupPumpLocked(*group); + pending_depth = group->pending_admissions.size(); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + queue_depth = group->queue.size(); + queued_bytes = group->queued_bytes; + active_sockets = activeSocketCountLocked(*group); +#endif + } + } + } + } catch (const std::exception& e) { + LOG(ERROR) << "Failed to admit TCP work for " << key.host << ":" + << key.port << ". Error: " << e.what(); + if (work.slice) rejected.emplace(std::move(work)); + rejection_reason = WorkFailureReason::RUNTIME_UNAVAILABLE; + } catch (...) { + LOG(ERROR) << "Failed to admit TCP work for " << key.host << ":" + << key.port << ". Error: unknown exception"; + if (work.slice) rejected.emplace(std::move(work)); + rejection_reason = WorkFailureReason::RUNTIME_UNAVAILABLE; + } + + cancelTimerNoThrow(timer_to_cancel); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (rejected) { + invokeLaneObserverHook(kLaneQueueRejected, queue_depth, queued_bytes, + active_sockets, false); + } else if (direct_admission) { + invokeLaneObserverHook(kLaneQueueAdmitted, queue_depth, queued_bytes, + active_sockets, false); + } + if (pending_admission) + invokeLaneObserverHook(kLaneAdmissionPending, pending_depth, 0, 0, + false); + if (promoted != 0) + invokeLaneObserverHook(kLaneAdmissionPromoted, pending_depth, promoted, + 0, false); + if (timer_armed) + invokeLaneObserverHook(kLaneAdmissionTimerArmed, pending_depth, 0, 0, + false); + if (hard_rejection) + invokeLaneObserverHook(kLaneAdmissionHardRejected, pending_depth, 0, 0, + false); +#endif + + failWorkItems(std::move(expired), WorkFailureReason::QUEUE_TIMEOUT, + state->failure_counters); + failWorkItems(std::move(runtime_failed), + WorkFailureReason::RUNTIME_UNAVAILABLE, + state->failure_counters); + if (rejected) + failWorkItem(std::move(*rejected), rejection_reason, + state->failure_counters); + else if (pump_epoch != 0) + postGroupPump(group, pump_epoch); +} + +void TcpTransport::postGroupPump( + const std::shared_ptr& group, uint64_t pump_epoch) { + auto fail_posted_work = [&](const char* error) { + std::deque failed_queue; + std::deque failed_pending; + std::shared_ptr admission_timer; + { + std::lock_guard lock(group->mutex); + if (group->pump_scheduled && group->pump_epoch == pump_epoch) { + group->pump_scheduled = false; + failed_queue.swap(group->queue); + clearQueuedBytesLocked(*group); + failed_pending.swap(group->pending_admissions); + ++group->admission_epoch; + if (group->admission_epoch == 0) ++group->admission_epoch; + admission_timer = std::move(group->admission_timer); + } + } + cancelTimerNoThrow(admission_timer); + failWorkItems(std::move(failed_queue), + WorkFailureReason::RUNTIME_UNAVAILABLE, + group->failure_counters); + failWorkItems(std::move(failed_pending), + WorkFailureReason::RUNTIME_UNAVAILABLE, + group->failure_counters); + LOG(ERROR) << "Failed to schedule TCP lane pump for " << group->key.host + << ":" << group->key.port + << (error && *error ? ". Error: " : "") + << (error && *error ? error : ""); + }; + + try { + asio::post(group->executor, + [group, pump_epoch] { runGroupPump(group, pump_epoch); }); + } catch (const std::exception& e) { + fail_posted_work(e.what()); + } catch (...) { + fail_posted_work(nullptr); + } +} + +void TcpTransport::runGroupPump( + const std::shared_ptr& group, uint64_t pump_epoch) { + struct LaneStart { + std::shared_ptr lane; + uint64_t epoch; + }; + std::array sessions; + std::array connects; + size_t session_count = 0; + size_t connect_count = 0; + std::deque failed; + std::deque expired; + std::deque runtime_failed; + std::shared_ptr admission_timer_to_cancel; + WorkFailureReason failure_reason = WorkFailureReason::CONNECT_FAILED; + uint64_t followup_pump_epoch = 0; + [[maybe_unused]] size_t promoted = 0; + [[maybe_unused]] size_t pending_depth = 0; + bool queue_detached_after_scheduling = false; + [[maybe_unused]] bool retry_armed = false; + [[maybe_unused]] bool cooldown_started = false; + [[maybe_unused]] bool admission_timer_armed = false; + + { + std::lock_guard lock(group->mutex); + if (!group->pump_scheduled || group->pump_epoch != pump_epoch) return; + group->pump_scheduled = false; + if (group->state != GroupState::OPEN) return; + + expirePendingAdmissionsLocked(*group, std::chrono::steady_clock::now(), + expired); + promoted += promotePendingAdmissionsLocked(*group); + + for (const auto& lane : group->lanes) { + if (group->queue.empty()) break; + if (lane->state != LaneState::IDLE) continue; + if (!lane->socket || !lane->socket->is_open()) { + lane->socket.reset(); + lane->state = LaneState::DISCONNECTED; + continue; + } + + lane->current.emplace(std::move(group->queue.front())); + const uint64_t length = lane->current->slice->length; + group->queue.pop_front(); + removeQueuedBytesLocked(*group, length); + if (group->queue.empty()) clearQueuedBytesLocked(*group); + promoted += promotePendingAdmissionsLocked(*group); + lane->state = LaneState::BUSY; + ++lane->operation_epoch; + if (lane->operation_epoch == 0) ++lane->operation_epoch; + sessions[session_count++] = {lane, lane->operation_epoch}; + } + + // Once armed, the retry timer exclusively owns the transition out of + // cooldown. A delayed pump must not reset round accounting or start a + // probe before the matching timer handler validates the group. + bool waiting_for_cooldown = group->retry_timer != nullptr; + const bool round_exhausted = hasDisconnectedLaneLocked(*group) && + !hasUntriedDisconnectedLaneLocked(*group); + if (!waiting_for_cooldown && !group->queue.empty() && + group->probes_in_flight == 0 && round_exhausted) { + const bool cooldown_already_started = + group->next_probe_not_before != + std::chrono::steady_clock::time_point{}; + if (!hasUsableLaneLocked(*group) && !cooldown_already_started) { + failed.swap(group->queue); + clearQueuedBytesLocked(*group); + enterReconnectCooldownLocked(*group); + failure_reason = WorkFailureReason::CONNECT_FAILED; + queue_detached_after_scheduling = true; + } else if (group->next_probe_not_before == + std::chrono::steady_clock::time_point{}) { + enterReconnectCooldownLocked(*group); + } + if (std::chrono::steady_clock::now() < + group->next_probe_not_before) { + if (armRetryTimerLocked(group)) { + waiting_for_cooldown = true; + retry_armed = group->retry_timer != nullptr; + } else { + failed.swap(group->queue); + clearQueuedBytesLocked(*group); + failure_reason = WorkFailureReason::RUNTIME_UNAVAILABLE; + queue_detached_after_scheduling = true; + } + } else { + beginConnectRoundLocked(*group); + } + } + + // A lane-local failure may defer only connection probes while healthy + // siblings continue pulling work above. Reuse the per-group timer so + // repeated failures cannot create a tight reconnect loop. + if (!waiting_for_cooldown && failed.empty() && !group->queue.empty() && + group->probes_in_flight == 0 && + std::chrono::steady_clock::now() < group->next_probe_not_before && + armRetryTimerLocked(group)) { + waiting_for_cooldown = group->retry_timer != nullptr; + retry_armed = waiting_for_cooldown; + } + + if (!waiting_for_cooldown && failed.empty()) { + const size_t probe_limit = + std::min(group->lanes.size(), kMaxConcurrentLaneProbes); + while (!group->queue.empty() && + group->probes_in_flight < probe_limit) { + auto lane_it = std::find_if( + group->lanes.begin(), group->lanes.end(), + [&group](const auto& lane) { + return lane->state == LaneState::DISCONNECTED && + lane->last_connect_round != group->connect_round; + }); + if (lane_it == group->lanes.end()) break; + + auto lane = *lane_it; + lane->state = LaneState::CONNECTING; + lane->connect_stage = LaneConnectStage::NONE; + lane->last_connect_round = group->connect_round; + ++lane->operation_epoch; + if (lane->operation_epoch == 0) ++lane->operation_epoch; + ++group->probes_in_flight; + connects[connect_count++] = {lane, lane->operation_epoch}; + } + + if (!group->queue.empty() && group->probes_in_flight == 0 && + hasDisconnectedLaneLocked(*group) && + !hasUntriedDisconnectedLaneLocked(*group)) { + if (!hasUsableLaneLocked(*group)) { + failed.swap(group->queue); + clearQueuedBytesLocked(*group); + enterReconnectCooldownLocked(*group); + queue_detached_after_scheduling = true; + } else { + enterReconnectCooldownLocked(*group); + } + cooldown_started = true; + } + } + + promoted += promotePendingAdmissionsLocked(*group); + refreshAdmissionTimerLocked(group, runtime_failed, + admission_timer_to_cancel, + admission_timer_armed); + pending_depth = group->pending_admissions.size(); + if (queue_detached_after_scheduling && !group->queue.empty()) + followup_pump_epoch = requestGroupPumpLocked(*group); + } + + cancelTimerNoThrow(admission_timer_to_cancel); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (retry_armed) invokeLaneObserverHook(kLaneRetryArmed, 0, 0, 0, false); + if (cooldown_started) + invokeLaneObserverHook(kLaneCooldownStarted, 0, 0, 0, false); + if (promoted != 0) + invokeLaneObserverHook(kLaneAdmissionPromoted, pending_depth, promoted, + 0, false); + if (admission_timer_armed) + invokeLaneObserverHook(kLaneAdmissionTimerArmed, pending_depth, 0, 0, + false); +#endif + for (size_t i = 0; i < connect_count; ++i) + startLaneConnect(group, connects[i].lane, connects[i].epoch); + for (size_t i = 0; i < session_count; ++i) + startLaneSession(group, sessions[i].lane, sessions[i].epoch); + failWorkItems(std::move(failed), failure_reason, group->failure_counters); + failWorkItems(std::move(expired), WorkFailureReason::QUEUE_TIMEOUT, + group->failure_counters); + failWorkItems(std::move(runtime_failed), + WorkFailureReason::RUNTIME_UNAVAILABLE, + group->failure_counters); + if (followup_pump_epoch != 0) postGroupPump(group, followup_pump_epoch); +} + +void TcpTransport::startLaneConnect( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch) { + std::string initiation_error; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + size_t queue_depth = 0; + uint64_t queued_bytes = 0; + size_t active_sockets = 0; +#endif + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + lane->state != LaneState::CONNECTING || + lane->operation_epoch != epoch) { + return; + } +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (invokeLaneConnectFailureInjectionHook(lane->lane_id)) { + initiation_error = "injected lane connect failure"; + } else +#endif + try { + lane->resolver = + std::make_shared(group->executor); + lane->socket = + std::make_shared(group->executor); + lane->connect_stage = LaneConnectStage::RESOLVING; + lane->resolver->async_resolve( + group->key.host, std::to_string(group->key.port), + [group, lane, epoch]( + asio::error_code ec, + asio::ip::tcp::resolver::results_type results) { +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneConnectHandlerHook(); +#endif + handleLaneResolved(group, lane, epoch, ec, + std::move(results)); + }); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + queue_depth = group->queue.size(); + queued_bytes = group->queued_bytes; + active_sockets = activeSocketCountLocked(*group); +#endif + } catch (const std::exception& e) { + initiation_error = e.what(); + } catch (...) { + initiation_error = "unknown exception"; + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (initiation_error.empty()) { + invokeLaneObserverHook(kLaneConnecting, queue_depth, queued_bytes, + active_sockets, false); + } +#endif + if (!initiation_error.empty()) + handleLaneConnectFailure(group, lane, epoch, initiation_error); +} + +void TcpTransport::handleLaneResolved( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch, + asio::error_code ec, asio::ip::tcp::resolver::results_type results) { + if (ec) { + handleLaneConnectFailure(group, lane, epoch, ec.message()); + return; + } + + std::string initiation_error; + [[maybe_unused]] bool stale = false; + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + lane->state != LaneState::CONNECTING || + lane->connect_stage != LaneConnectStage::RESOLVING || + lane->operation_epoch != epoch || !lane->socket) { + stale = true; + } else { + lane->connect_stage = LaneConnectStage::CONNECTING; + try { + asio::async_connect( + *lane->socket, results, + [group, lane, epoch](asio::error_code connect_ec, + const asio::ip::tcp::endpoint&) { + handleLaneConnected(group, lane, epoch, connect_ec); + }); + } catch (const std::exception& e) { + initiation_error = e.what(); + } catch (...) { + initiation_error = "unknown exception"; + } + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (stale) invokeLaneObserverHook(kLaneLateHandler, 0, 0, 0, false); +#endif + if (!initiation_error.empty()) + handleLaneConnectFailure(group, lane, epoch, initiation_error); +} + +void TcpTransport::handleLaneConnected( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch, + asio::error_code ec) { + if (ec) { + handleLaneConnectFailure(group, lane, epoch, ec.message()); + return; + } + + uint64_t pump_epoch = 0; + [[maybe_unused]] bool stale = false; + std::string option_error; + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + lane->state != LaneState::CONNECTING || + lane->connect_stage != LaneConnectStage::CONNECTING || + lane->operation_epoch != epoch || !lane->socket) { + stale = true; + } else { + asio::error_code option_ec; + lane->socket->set_option(asio::ip::tcp::no_delay(true), option_ec); + if (option_ec) { + option_error = option_ec.message(); + } else { + if (group->probes_in_flight != 0) --group->probes_in_flight; + group->connect_round_had_success = true; + lane->resolver.reset(); + lane->connect_stage = LaneConnectStage::NONE; + lane->state = LaneState::IDLE; + pump_epoch = requestGroupPumpLocked(*group); + } + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (stale) invokeLaneObserverHook(kLaneLateHandler, 0, 0, 0, false); +#endif + if (!option_error.empty()) { + handleLaneConnectFailure(group, lane, epoch, option_error); + } else if (pump_epoch != 0) { + postGroupPump(group, pump_epoch); + } +} + +void TcpTransport::handleLaneConnectFailure( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch, + const std::string& error) { + std::shared_ptr resolver; + std::shared_ptr socket; + std::deque failed; + std::deque expired; + std::deque runtime_failed; + std::shared_ptr admission_timer_to_cancel; + uint64_t pump_epoch = 0; + [[maybe_unused]] size_t promoted = 0; + [[maybe_unused]] size_t pending_depth = 0; + bool stale = false; + [[maybe_unused]] bool cooldown_started = false; + [[maybe_unused]] bool admission_timer_armed = false; + uint64_t connect_failure_log_count = 0; + { + std::lock_guard lock(group->mutex); + if (lane->state != LaneState::CONNECTING || + lane->operation_epoch != epoch) { + stale = true; + } else { + connect_failure_log_count = ++group->connect_failure_log_count; + if (group->probes_in_flight != 0) --group->probes_in_flight; + resolver = std::move(lane->resolver); + socket = std::move(lane->socket); + lane->connect_stage = LaneConnectStage::NONE; + lane->state = group->state == GroupState::OPEN + ? LaneState::DISCONNECTED + : LaneState::CLOSING; + + const bool sibling_usable = + group->state == GroupState::OPEN && hasUsableLaneLocked(*group); + + if (group->state == GroupState::OPEN && !group->queue.empty() && + !sibling_usable && group->probes_in_flight == 0 && + !hasUntriedDisconnectedLaneLocked(*group)) { + failed.swap(group->queue); + clearQueuedBytesLocked(*group); + enterReconnectCooldownLocked(*group); + cooldown_started = true; + expirePendingAdmissionsLocked( + *group, std::chrono::steady_clock::now(), expired); + promoted = promotePendingAdmissionsLocked(*group); + refreshAdmissionTimerLocked(group, runtime_failed, + admission_timer_to_cancel, + admission_timer_armed); + pending_depth = group->pending_admissions.size(); + pump_epoch = requestGroupPumpLocked(*group); + } else { + pump_epoch = requestGroupPumpLocked(*group); + } + } + } + + if (resolver) { + try { + resolver->cancel(); + } catch (...) { + } + } + closeSocketNoThrow(socket); + cancelTimerNoThrow(admission_timer_to_cancel); + if (!stale && shouldLogOccurrence(connect_failure_log_count)) { + LOG(ERROR) << "TCP lane connection to " << group->key.host << ":" + << group->key.port << " failed: " << error + << " (attempt failure " << connect_failure_log_count << ")"; + } +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (stale) invokeLaneObserverHook(kLaneLateHandler, 0, 0, 0, false); + if (cooldown_started) + invokeLaneObserverHook(kLaneCooldownStarted, 0, 0, 0, false); + if (promoted != 0) + invokeLaneObserverHook(kLaneAdmissionPromoted, pending_depth, promoted, + 0, false); + if (admission_timer_armed) + invokeLaneObserverHook(kLaneAdmissionTimerArmed, pending_depth, 0, 0, + false); +#endif + failWorkItems(std::move(failed), WorkFailureReason::CONNECT_FAILED, + group->failure_counters); + failWorkItems(std::move(expired), WorkFailureReason::QUEUE_TIMEOUT, + group->failure_counters); + failWorkItems(std::move(runtime_failed), + WorkFailureReason::RUNTIME_UNAVAILABLE, + group->failure_counters); + if (pump_epoch != 0) postGroupPump(group, pump_epoch); +} + +void TcpTransport::startLaneSession( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch) { + std::string initiation_error; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + size_t queue_depth = 0; + uint64_t queued_bytes = 0; + size_t active_sockets = 0; +#endif + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + lane->state != LaneState::BUSY || lane->operation_epoch != epoch || + !lane->current) { + return; + } + if (!lane->socket || !lane->socket->is_open()) { + initiation_error = "lane socket is not open"; + } else { + // This function runs on the TCP executor. Keep construction and + // initial Asio initiation under the group lock so shutdown cannot + // invalidate the checked epoch between validation and initiation. + // Asio initiating functions do not invoke their completion handler + // inline, so this cannot call lane terminal completion under the + // mutex. + try { + std::weak_ptr weak_group(group); + std::weak_ptr weak_lane(lane); + auto session = std::make_shared( + lane->socket, lane->current->use_v2, + [weak_group, weak_lane, epoch](TransferStatusEnum status, + bool clean) noexcept { + auto callback_group = weak_group.lock(); + auto callback_lane = weak_lane.lock(); + if (!callback_group || !callback_lane) return; + handleLaneTerminal(callback_group, callback_lane, epoch, + status, clean); + }); + lane->session = session; + session->initiate(lane->current->slice->source_addr, + lane->current->slice->tcp.dest_addr, + lane->current->slice->length, + lane->current->slice->opcode); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + queue_depth = group->queue.size(); + queued_bytes = group->queued_bytes; + active_sockets = activeSocketCountLocked(*group); +#endif + } catch (const std::exception& e) { + lane->session.reset(); + initiation_error = e.what(); + } catch (...) { + lane->session.reset(); + initiation_error = "unknown exception"; + } + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (initiation_error.empty()) { + invokeLaneObserverHook(kLaneBusy, queue_depth, queued_bytes, + active_sockets, true); + } +#endif + if (!initiation_error.empty()) { + LOG(ERROR) << "Failed to start TCP lane session for " << group->key.host + << ":" << group->key.port << ". Error: " << initiation_error; + handleLaneTerminal(group, lane, epoch, TransferStatusEnum::FAILED, + false); + } +} + +void TcpTransport::handleLaneTerminal( + const std::shared_ptr& group, + const std::shared_ptr& lane, uint64_t epoch, + TransferStatusEnum status, bool connection_clean) noexcept { + std::optional action; + std::shared_ptr socket_to_close; + bool stale = false; + { + std::lock_guard lock(group->mutex); + if (lane->operation_epoch != epoch || lane->state != LaneState::BUSY || + !lane->current) { + stale = true; + } else { + action.emplace(std::move(*lane->current), status, connection_clean); + lane->current.reset(); + lane->session.reset(); + lane->state = LaneState::COMPLETING; + if (!connection_clean || group->state != GroupState::OPEN) + socket_to_close = std::move(lane->socket); + } + } + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + if (stale) { + invokeLaneObserverHook(kLaneLateHandler, 0, 0, 0, false); + } +#endif + if (stale) return; + + if (status != TransferStatusEnum::COMPLETED) { + recordWorkFailure(WorkFailureReason::SESSION_FAILED, + group->failure_counters); + } + + // A dirty protocol stream must be closed before terminal Slice status is + // visible to the caller. + closeSocketNoThrow(socket_to_close); + completeTerminalAction(std::move(*action)); + + uint64_t pump_epoch = 0; + std::shared_ptr shutdown_socket; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + size_t queue_depth = 0; + uint64_t queued_bytes = 0; + size_t active_sockets = 0; +#endif + { + std::lock_guard lock(group->mutex); + if (group->state != GroupState::OPEN || + lane->operation_epoch != epoch) { + lane->state = LaneState::CLOSING; + shutdown_socket = std::move(lane->socket); + } else if (connection_clean && lane->socket && + lane->socket->is_open()) { + lane->state = LaneState::IDLE; + } else { + lane->socket.reset(); + lane->state = LaneState::DISCONNECTED; + // Keep this lane marked as tried in the current round. If another + // lane is still usable, wait for the group cooldown before a new + // round can retry disconnected lanes. If this was the last usable + // lane and the round had previously connected successfully, start + // a fresh round immediately so queued work is not mistaken for an + // all-probes-failed round. + const bool sibling_usable = hasUsableLaneLocked(*group); + if (sibling_usable) { + enterReconnectCooldownLocked(*group); + } else if (!group->queue.empty() && group->probes_in_flight == 0 && + group->connect_round_had_success) { + beginConnectRoundLocked(*group); + } + } + pump_epoch = requestGroupPumpLocked(*group); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + queue_depth = group->queue.size(); + queued_bytes = group->queued_bytes; + active_sockets = activeSocketCountLocked(*group); +#endif + } + closeSocketNoThrow(shutdown_socket); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneObserverHook(kLaneTerminal, queue_depth, queued_bytes, + active_sockets, false); +#endif + if (pump_epoch != 0) postGroupPump(group, pump_epoch); +} + +void TcpTransport::completeTerminalAction(TerminalAction action) noexcept { + auto continuation = std::move(action.work.continuation); + try { + if (action.status == TransferStatusEnum::COMPLETED) + action.work.slice->markSuccess(); + else + action.work.slice->markFailed(); + } catch (const std::exception& e) { + LOG(ERROR) << "TCP Slice terminal completion threw: " << e.what(); + } catch (...) { + LOG(ERROR) << "TCP Slice terminal completion threw"; + } + + if (continuation) { + try { + continuation(); + } catch (const std::exception& e) { + LOG(ERROR) << "TCP Slice continuation threw: " << e.what(); + } catch (...) { + LOG(ERROR) << "TCP Slice continuation threw"; + } + } +} + +uint64_t TcpTransport::recordWorkFailure( + WorkFailureReason reason, + const std::shared_ptr& counters) noexcept { + if (!counters) return 0; + + std::atomic* counter = nullptr; + switch (reason) { + case WorkFailureReason::QUEUE_FULL: + counter = &counters->queue_full; + break; + case WorkFailureReason::QUEUE_TIMEOUT: + counter = &counters->queue_timeout; + break; + case WorkFailureReason::RUNTIME_UNAVAILABLE: + counter = &counters->runtime_unavailable; + break; + case WorkFailureReason::CONNECT_FAILED: + counter = &counters->connect_failed; + break; + case WorkFailureReason::SESSION_FAILED: + counter = &counters->session_failed; + break; + case WorkFailureReason::SHUTDOWN: + counter = &counters->shutdown; + break; + } + if (!counter) return 0; + + const uint64_t occurrence = + counter->fetch_add(1, std::memory_order_relaxed) + 1; + if (reason == WorkFailureReason::QUEUE_FULL && + shouldLogOccurrence(occurrence)) { + LOG(WARNING) << "TCP lane queue-full rejection count: " << occurrence; + } +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneFailureReasonHook(static_cast(reason)); +#endif + return occurrence; +} + +void TcpTransport::failWorkItem( + TcpWorkItem work, WorkFailureReason reason, + const std::shared_ptr& counters) noexcept { + recordWorkFailure(reason, counters); + completeTerminalAction( + TerminalAction(std::move(work), TransferStatusEnum::FAILED, false)); +} + +void TcpTransport::failWorkItems( + std::deque work, WorkFailureReason reason, + const std::shared_ptr& counters) noexcept { + while (!work.empty()) { + auto item = std::move(work.front()); + work.pop_front(); + failWorkItem(std::move(item), reason, counters); + } +} + +void TcpTransport::closeSocketNoThrow( + const std::shared_ptr& socket) noexcept { + if (!socket) return; + asio::error_code error; + socket->cancel(error); + socket->close(error); +} + +void TcpTransport::shutdownConnectionLanes() { + const auto state = lane_state_; + std::vector> groups; + { + std::lock_guard state_lock(state->mutex); + if (state->shutting_down) return; + state->shutting_down = true; + groups.reserve(state->groups.size()); + for (const auto& entry : state->groups) groups.push_back(entry.second); + } + + for (const auto& group : groups) { + std::deque accepted_queue; + std::deque accepted_pending; + { + std::lock_guard lock(group->mutex); + group->state = GroupState::CLOSING; + group->pump_scheduled = false; + ++group->pump_epoch; + accepted_queue.swap(group->queue); + accepted_pending.swap(group->pending_admissions); + clearQueuedBytesLocked(*group); + ++group->retry_epoch; + if (group->retry_epoch == 0) ++group->retry_epoch; + ++group->admission_epoch; + if (group->admission_epoch == 0) ++group->admission_epoch; + for (const auto& lane : group->lanes) { + ++lane->operation_epoch; + if (lane->operation_epoch == 0) ++lane->operation_epoch; + if (lane->state != LaneState::CLOSED) + lane->state = LaneState::CLOSING; + } + } + failWorkItems(std::move(accepted_queue), WorkFailureReason::SHUTDOWN, + state->failure_counters); + failWorkItems(std::move(accepted_pending), WorkFailureReason::SHUTDOWN, + state->failure_counters); + } + + auto cancellation_posts = std::make_shared(); + if (context_ && running_) { + for (const auto& group : groups) { + cancellation_posts->add(); + try { + asio::post(group->executor, [group, cancellation_posts] { + std::vector> sessions; + std::vector> + resolvers; + std::vector> sockets; + std::shared_ptr retry_timer; + std::shared_ptr admission_timer; + { + std::lock_guard lock(group->mutex); + retry_timer = group->retry_timer; + admission_timer = group->admission_timer; + for (const auto& lane : group->lanes) { + if (lane->session) + sessions.push_back(lane->session); + if (lane->resolver) + resolvers.push_back(lane->resolver); + if (lane->socket) sockets.push_back(lane->socket); + } + } + for (const auto& session : sessions) + if (session) session->cancel(); + for (const auto& resolver : resolvers) { + if (!resolver) continue; + try { + resolver->cancel(); + } catch (...) { + } + } + for (const auto& socket : sockets) + closeSocketNoThrow(socket); + if (retry_timer) { + asio::error_code timer_ec; + retry_timer->cancel(timer_ec); + } + cancelTimerNoThrow(admission_timer); + cancellation_posts->done(); + }); + } catch (...) { + cancellation_posts->done(); + } + } + cancellation_posts->waitUntil(std::chrono::steady_clock::now() + + kShutdownCancellationWait); + } + + running_ = false; + if (context_) context_->io_context.stop(); + if (thread_.joinable()) thread_.join(); + + std::deque deferred; + std::vector> sessions; + std::vector> resolvers; + std::vector> sockets; + std::vector> retry_timers; + std::vector> admission_timers; + + for (const auto& group : groups) { + { + std::lock_guard lock(group->mutex); + if (group->retry_timer) + retry_timers.push_back(std::move(group->retry_timer)); + if (group->admission_timer) + admission_timers.push_back(std::move(group->admission_timer)); + for (const auto& lane : group->lanes) { + if (lane->current) { + deferred.emplace_back(std::move(*lane->current)); + lane->current.reset(); + } + if (lane->session) sessions.push_back(std::move(lane->session)); + if (lane->resolver) + resolvers.push_back(std::move(lane->resolver)); + if (lane->socket) sockets.push_back(std::move(lane->socket)); + lane->connect_stage = LaneConnectStage::NONE; + lane->state = LaneState::CLOSED; + } + group->probes_in_flight = 0; + group->state = GroupState::CLOSED; + } +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + invokeLaneObserverHook(kLaneShutdownClean, 0, 0, 0, false); +#endif + } + + // No handler is running after join. Reset every Asio-owning field while + // TcpContext and its execution_context are still alive, then publish + // terminal failure for work that had been BUSY. + for (const auto& session : sessions) + if (session) session->cancel(); + for (const auto& resolver : resolvers) { + if (!resolver) continue; + try { + resolver->cancel(); + } catch (...) { + } + } + for (const auto& socket : sockets) closeSocketNoThrow(socket); + for (const auto& timer : retry_timers) { + cancelTimerNoThrow(timer); + } + for (const auto& timer : admission_timers) cancelTimerNoThrow(timer); + sessions.clear(); + resolvers.clear(); + sockets.clear(); + retry_timers.clear(); + admission_timers.clear(); + + failWorkItems(std::move(deferred), WorkFailureReason::SHUTDOWN, + state->failure_counters); + + const auto& counters = state->failure_counters; + VLOG(1) << "TCP lane failure totals: queue_full=" + << counters->queue_full.load(std::memory_order_relaxed) + << ", queue_timeout=" + << counters->queue_timeout.load(std::memory_order_relaxed) + << ", connect_failed=" + << counters->connect_failed.load(std::memory_order_relaxed) + << ", runtime_unavailable=" + << counters->runtime_unavailable.load(std::memory_order_relaxed) + << ", session_failed=" + << counters->session_failed.load(std::memory_order_relaxed) + << ", shutdown=" + << counters->shutdown.load(std::memory_order_relaxed); + + { + std::lock_guard state_lock(state->mutex); + state->groups.clear(); + state->runtime.reset(); + } + groups.clear(); + lane_runtime_.reset(); +} + +bool TcpTransport::validateAddress(uint64_t addr, uint64_t size) const { + return validateTcpAddress(metadata_, addr, size); +} diff --git a/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_session_impl.h b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_session_impl.h new file mode 100644 index 0000000000..4cb8907e48 --- /dev/null +++ b/mooncake-transfer-engine/src/transport/tcp_transport/tcp_transport_session_impl.h @@ -0,0 +1,815 @@ +// Copyright 2024 KVCache.AI +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Private implementation fragment included by tcp_transport.cpp from inside +// namespace mooncake. Keeping the protocol/session state here makes the public +// transport entry points and the lane scheduler independently reviewable +// without changing linkage or lifetime behavior. + +static size_t getChunkSize() { + static const size_t val = [] { + const char* env = std::getenv("MC_TCP_SLICE_SIZE"); + if (env) { + try { + size_t v = std::stoull(env); + if (v > 0) return v; + LOG(WARNING) + << "Ignore non-positive MC_TCP_SLICE_SIZE value: " << env + << ", using default 65536"; + } catch (const std::exception& e) { + // A non-numeric or out-of-range value makes std::stoull throw; + // fall through to the default instead of letting the exception + // propagate out of this static initializer and abort the + // transfer that first reads the chunk size. + LOG(WARNING) + << "Invalid MC_TCP_SLICE_SIZE value: " << env + << ". Error: " << e.what() << ", using default 65536"; + } + } + return size_t(65536); // 64KB default + }(); + return val; +} + +struct SessionHeader { + uint64_t size; + uint64_t addr; + uint8_t opcode; +}; + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) +// Returns the CUDA device ordinal if addr is device memory, or -1 otherwise. +// Callers must call cudaSetDevice before any cudaMemcpy to avoid implicit +// GPU 0 context creation. +static int getCudaDeviceId(void* addr) { + cudaPointerAttributes attributes; + auto status = cudaPointerGetAttributes(&attributes, addr); + if (status != cudaSuccess) return -1; + if (attributes.type == cudaMemoryTypeDevice) return attributes.device; + return -1; +} + +#ifdef USE_MACA +static cudaError_t copyTcpCudaMemory(void* dst, const void* src, size_t size) { + cudaStream_t stream; + cudaError_t status = + cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking); + if (status != cudaSuccess) return status; + + status = cudaMemcpyAsync(dst, src, size, cudaMemcpyDefault, stream); + if (status == cudaSuccess) { + status = cudaStreamSynchronize(stream); + } + + cudaError_t destroy_status = cudaStreamDestroy(stream); + return status == cudaSuccess ? destroy_status : status; +} +#endif +#endif + +// Forward declaration +class TcpTransport; + +using ValidateAddrFn = std::function; + +// --- Acknowledged framing (protocol v2, #2086) ------------------------------ +// v1 framing gives the initiator no channel to learn whether the receiver +// applied (or even accepted) a WRITE: COMPLETED fires when the final chunk +// enters the initiator's kernel socket buffer, while megabytes may still be +// in flight toward destination memory, and a rejected request is silently +// "successful". v2 requests set the high bit of the opcode; the server then +// (a) prefixes every READ response with an 8-byte status frame and (b) sends +// an 8-byte status frame for WRITE only after the final chunk has been +// applied to destination memory. Initiators enable v2 only when the target +// segment advertises tcp_proto_version >= 2, so old servers never see +// flagged opcodes and old initiators keep receiving v1 framing. +static constexpr uint8_t kOpcodeV2Flag = 0x80; +// Status frames carry a magic in the high 32 bits so that a v2 initiator +// which reaches a v1 server through a stale descriptor (v1 treats unknown +// opcodes as READ and immediately streams payload bytes) fails fast on the +// first frame instead of misinterpreting the stream. Residual risk: payload +// bytes that happen to equal a valid frame (2^-64 per request, data +// dependent) are indistinguishable in-band; eliminating that would need a +// nonce/checksum handshake, which this deliberately avoids. +static constexpr uint64_t kStatusMagic = 0x4D435456ull << 32; // "MCTV" +static constexpr uint64_t kStatusOk = kStatusMagic | 0; +static constexpr uint64_t kStatusAddrRejected = kStatusMagic | 1; +static inline bool statusFrameValid(uint64_t frame) { + return (frame & 0xFFFFFFFF00000000ull) == kStatusMagic; +} + +// Operational escape hatch: MC_TCP_PROTO=1 forces initiators to speak the +// legacy unacknowledged framing even to v2-capable servers. Also used by +// tests to cover the mixed-version matrix in one process. +static bool forceLegacyTcpProto() { + // Read per call (startTransfer already does metadata lookups; getenv is + // noise) so tests can cover both protocol modes in one process. + const char* env = std::getenv("MC_TCP_PROTO"); + return env && env[0] == '1' && env[1] == '\0'; +} + +// Server-side session: handles transfer requests on a persistent connection. +// The session owns the socket; ending the callback chain without rearming +// (start()/next handler) drops the last reference and closes the connection. +struct ServerSession : public std::enable_shared_from_this { + explicit ServerSession(std::shared_ptr socket, + ValidateAddrFn validate_addr) + : socket_(std::move(socket)), + validate_addr_(std::move(validate_addr)) {} + + std::shared_ptr socket_; + ValidateAddrFn validate_addr_; + SessionHeader header_; + uint64_t total_transferred_bytes_; + char* local_buffer_; + bool v2_ = false; + uint64_t status_frame_; + + void start() { + total_transferred_bytes_ = 0; + readHeader(); + } + + private: + // Send an 8-byte status frame, then run `next` (or end the session — + // closing the connection — when `next` is empty or the send fails). + void sendStatus(uint64_t status, std::function next) { + status_frame_ = htole64(status); + auto self(shared_from_this()); + asio::async_write(*socket_, + asio::buffer(&status_frame_, sizeof(status_frame_)), + [this, self, next = std::move(next)]( + const asio::error_code& ec, std::size_t) { + if (ec) + return; // connection closes with the session + if (next) next(); + }); + } + + void readHeader() { + auto self(shared_from_this()); + asio::async_read( + *socket_, asio::buffer(&header_, sizeof(SessionHeader)), + [this, self](const asio::error_code& ec, std::size_t len) { + if (ec || len != sizeof(SessionHeader)) { + if (ec.value() != asio::error::eof) { + LOG(WARNING) + << "ServerSession::readHeader failed. Error: " + << ec.message() << " (value: " << ec.value() << ")" + << ", bytes read: " << len; + } + return; + } + + v2_ = (header_.opcode & kOpcodeV2Flag) != 0; + const uint8_t opcode = header_.opcode & ~kOpcodeV2Flag; + local_buffer_ = (char*)(le64toh(header_.addr)); + uint64_t size = le64toh(header_.size); + if (validate_addr_ && + !validate_addr_((uint64_t)local_buffer_, size)) { + LOG(ERROR) << "ServerSession: remote-supplied address 0x" + << std::hex << (uint64_t)local_buffer_ + << std::dec << " with size " << size + << " is not within any registered buffer"; + // v2 initiators learn of the rejection; v1 initiators + // only see the connection close (and, for small WRITEs, + // may have already reported success — the defect v2 + // exists to fix). + if (v2_) sendStatus(kStatusAddrRejected, nullptr); + return; + } + if (opcode == (uint8_t)TransferRequest::WRITE) { + readBody(); + } else if (v2_) { + // READ, v2: status frame precedes the data. + sendStatus(kStatusOk, [this] { writeBody(); }); + } else { + writeBody(); + } + }); + } + + void writeBody() { + auto self(shared_from_this()); + uint64_t size = le64toh(header_.size); + char* addr = local_buffer_; + + size_t buffer_size = + std::min(getChunkSize(), size - total_transferred_bytes_); + if (buffer_size == 0) { + // Transfer complete, wait for next request on this connection + start(); + return; + } + + char* dram_buffer = addr + total_transferred_bytes_; + int cuda_device = -1; + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + cuda_device = getCudaDeviceId(addr); + if (cuda_device >= 0) { + dram_buffer = new char[buffer_size]; + cudaSetDevice(cuda_device); +#ifdef USE_MACA + cudaError_t cuda_status = copyTcpCudaMemory( + dram_buffer, addr + total_transferred_bytes_, buffer_size); +#else + cudaError_t cuda_status = + cudaMemcpy(dram_buffer, addr + total_transferred_bytes_, + buffer_size, cudaMemcpyDefault); +#endif + if (cuda_status != cudaSuccess) { + LOG(ERROR) << "ServerSession::writeBody failed to copy from " + "CUDA memory. " + << "Error: " << cudaGetErrorString(cuda_status); + delete[] dram_buffer; + return; // Connection will be closed + } + } +#endif + + asio::async_write( + *socket_, asio::buffer(dram_buffer, buffer_size), + [this, addr, dram_buffer, cuda_device, self]( + const asio::error_code& ec, std::size_t transferred_bytes) { +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + if (cuda_device >= 0) { + delete[] dram_buffer; + } +#endif + if (ec) { + LOG(ERROR) + << "ServerSession::writeBody failed. " + << "Attempt to write data " << static_cast(addr) + << " using buffer " << static_cast(dram_buffer) + << ". Error: " << ec.message() + << " (value: " << ec.value() << ")"; + return; // Connection will be closed + } + total_transferred_bytes_ += transferred_bytes; + writeBody(); + }); + } + + void readBody() { + auto self(shared_from_this()); + uint64_t size = le64toh(header_.size); + char* addr = local_buffer_; + + size_t buffer_size = + std::min(getChunkSize(), size - total_transferred_bytes_); + if (buffer_size == 0) { + // Destination memory now holds the complete payload. Under v2, + // acknowledge before accepting the next request — this is what + // makes the initiator's COMPLETED mean "applied at the + // destination" rather than "left my socket buffer". + if (v2_) { + sendStatus(kStatusOk, [this] { start(); }); + } else { + start(); + } + return; + } + + char* dram_buffer = addr + total_transferred_bytes_; + int cuda_device = -1; + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + cuda_device = getCudaDeviceId(addr); + if (cuda_device >= 0) { + dram_buffer = new char[buffer_size]; + } +#endif + + asio::async_read( + *socket_, asio::buffer(dram_buffer, buffer_size), + [this, addr, dram_buffer, cuda_device, self]( + const asio::error_code& ec, std::size_t transferred_bytes) { + if (ec) { + // If client closed connection (EOF), this is normal - don't + // log + if (ec.value() != asio::error::eof) { + LOG(WARNING) + << "ServerSession::readBody failed. " + << "Attempt to read data " + << static_cast(addr) << " using buffer " + << static_cast(dram_buffer) + << ". Error: " << ec.message() + << " (value: " << ec.value() << ")"; + } + if (cuda_device >= 0) delete[] dram_buffer; + return; // Connection will be closed + } + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + if (cuda_device >= 0) { + cudaSetDevice(cuda_device); +#ifdef USE_MACA + cudaError_t cuda_status = + copyTcpCudaMemory(addr + total_transferred_bytes_, + dram_buffer, transferred_bytes); +#else + cudaError_t cuda_status = + cudaMemcpy(addr + total_transferred_bytes_, dram_buffer, + transferred_bytes, cudaMemcpyDefault); +#endif + if (cuda_status != cudaSuccess) { + LOG(ERROR) + << "ServerSession::readBody failed to copy to CUDA " + "memory. " + << "Error: " << cudaGetErrorString(cuda_status); + delete[] dram_buffer; + return; // Connection will be closed + } + delete[] dram_buffer; + } +#endif + total_transferred_bytes_ += transferred_bytes; + readBody(); + }); + } +}; + +// Client-side session: initiates one transfer request +struct ClientSession : public std::enable_shared_from_this { + using OnTerminal = + std::function; + + explicit ClientSession(std::shared_ptr socket, bool use_v2, + OnTerminal on_terminal) + : socket_(std::move(socket)), + v2_(use_v2), + on_terminal_(std::move(on_terminal)) {} + + std::shared_ptr socket_; + SessionHeader header_; + uint64_t total_transferred_bytes_; + char* local_buffer_; + bool v2_; + uint64_t status_frame_; + // v2 WRITE runs the body stream and the ack read concurrently (one + // async op per direction; handlers serialize on the io thread). The + // concurrent read lets a rejection — or a v1 server's bogus payload — + // abort a large in-flight WRITE instead of deadlocking on mutually + // full socket buffers, and delivers rejection frames before the close. + bool write_body_done_ = false; + bool write_acked_ok_ = false; + // An early negative/malformed ack can arrive while asio::async_write still + // owns a buffer pointing into the caller's source memory. Do not publish a + // terminal status until that body operation has completed or been + // cancelled: callers are allowed to release the source buffer as soon as + // the transfer becomes terminal. + bool write_body_in_flight_ = false; + bool write_abort_requested_ = false; + // A v2 status frame is prompt by construction: a READ status precedes + // any payload, and a WRITE ack follows at most one chunk's apply after + // the body is done. The only peer that never sends one is a legacy + // server reached through a stale v2 descriptor — and for requests + // shorter than a frame it also keeps the connection open (it streamed + // size < 8 bytes of "READ payload" and is waiting for our next header), + // so without a deadline both sides wait forever. Bound that wait; the + // default is generous so no healthy slow path can trip it, since it only + // covers the frame itself, never payload streaming. + static int statusFrameTimeoutSec() { + const char* env = std::getenv("MC_TCP_STATUS_TIMEOUT_SEC"); + if (env) { + int v = std::atoi(env); + if (v > 0) return v; + } + return 30; + } + std::optional status_timer_; + bool status_deadline_disarmed_ = false; + bool terminal_reported_ = false; + OnTerminal on_terminal_; + + void initiate(void* buffer, uint64_t dest_addr, size_t size, + TransferRequest::OpCode opcode) { + local_buffer_ = (char*)buffer; + header_.addr = htole64(dest_addr); + header_.size = htole64(size); + header_.opcode = (uint8_t)opcode | (v2_ ? kOpcodeV2Flag : 0); + total_transferred_bytes_ = 0; + writeHeader(); + } + + void cancel() noexcept { + cancelStatusDeadline(); + if (!socket_) return; + asio::error_code cancel_ec; + socket_->cancel(cancel_ec); + asio::error_code close_ec; + socket_->close(close_ec); + } + + private: + // All handlers run on the transport's single io thread, so arm/cancel + // and the expiry handler never race. Expiry only closes the socket: the + // pending status read then completes with an error and its handler owns + // the failure path (including source-buffer quiescence for WRITE). + void armStatusDeadline() { + auto self(shared_from_this()); + status_deadline_disarmed_ = false; + status_timer_.emplace(socket_->get_executor()); + status_timer_->expires_after( + std::chrono::seconds(statusFrameTimeoutSec())); + status_timer_->async_wait([this, self](const asio::error_code& ec) { + // The disarmed flag also covers an expiry that was already + // queued when cancel() ran (cancel cannot revoke those, and by + // then the socket may have been re-pooled). + if (ec == asio::error::operation_aborted || + status_deadline_disarmed_) { + return; + } + LOG(ERROR) << "ClientSession: no status frame within " + << statusFrameTimeoutSec() + << "s (peer likely speaks the legacy protocol); " + "dropping connection"; + if (socket_ && socket_->is_open()) { + asio::error_code cec; + socket_->close(cec); + } + }); + } + + void cancelStatusDeadline() { + status_deadline_disarmed_ = true; + if (status_timer_) { + asio::error_code ec; + status_timer_->cancel(ec); + } + } + + // Single terminal path. The invoking Asio operation has already released + // its buffer before entering its completion handler. The lane posts any + // follow-up pump, so a clean socket cannot be reused inline here. + void finalize(TransferStatusEnum status, bool clean) { + if (terminal_reported_) return; + terminal_reported_ = true; + cancelStatusDeadline(); + auto on_terminal = std::move(on_terminal_); + if (on_terminal) on_terminal(status, clean); + } + + // Abort a v2 WRITE and cancel any body operation. If asio still owns the + // current source buffer, its completion handler is responsible for + // finalizing after the buffer is quiescent. + void abortWrite() { + write_abort_requested_ = true; + if (socket_ && socket_->is_open()) { + asio::error_code ec; + socket_->close(ec); + } + if (!write_body_in_flight_) finalize(TransferStatusEnum::FAILED, false); + } + + void writeHeader() { + auto self(shared_from_this()); + asio::async_write( + *socket_, asio::buffer(&header_, sizeof(SessionHeader)), + [this, self](const asio::error_code& ec, std::size_t len) { + if (ec || len != sizeof(SessionHeader)) { + LOG(ERROR) + << "ClientSession::writeHeader failed. Error: " + << ec.message() << " (value: " << ec.value() << ")" + << ", bytes written: " << len; + finalize(TransferStatusEnum::FAILED, false); + return; + } + if ((header_.opcode & ~kOpcodeV2Flag) == + (uint8_t)TransferRequest::WRITE) { + if (v2_) readWriteAck(); // concurrent with the body + writeBody(); + } else if (v2_) { + readReadStatus(); + } else { + readBody(); + } + }); + } + + // v2 READ: the server prefixes the data with a status frame. + void readReadStatus() { + auto self(shared_from_this()); + armStatusDeadline(); + asio::async_read( + *socket_, asio::buffer(&status_frame_, sizeof(status_frame_)), + [this, self](const asio::error_code& ec, std::size_t len) { + cancelStatusDeadline(); + if (ec || len != sizeof(status_frame_)) { + LOG(ERROR) + << "ClientSession: failed to read READ status " + "frame. Error: " + << ec.message() << " (value: " << ec.value() << ")"; + finalize(TransferStatusEnum::FAILED, false); + return; + } + uint64_t frame = le64toh(status_frame_); + if (!statusFrameValid(frame)) { + LOG(ERROR) << "ClientSession: malformed READ status " + "frame (peer likely speaks the legacy " + "protocol); dropping connection"; + finalize(TransferStatusEnum::FAILED, false); + return; + } + if (frame != kStatusOk) { + LOG(ERROR) << "ClientSession: READ rejected by server, " + "status " + << (frame & 0xFFFFFFFFull); + finalize(TransferStatusEnum::FAILED, false); + return; + } + readBody(); + }); + } + + // v2 WRITE: completion is the server's acknowledgment that the payload + // has been applied to destination memory. Armed concurrently with the + // body stream; a well-behaved v2 server only sends the frame after the + // final chunk, so a frame arriving before the body is done is either a + // rejection or a legacy peer's payload — both close the socket immediately + // and publish failure only after the outstanding body write is quiescent. + void readWriteAck() { + auto self(shared_from_this()); + asio::async_read( + *socket_, asio::buffer(&status_frame_, sizeof(status_frame_)), + [this, self](const asio::error_code& ec, std::size_t len) { + cancelStatusDeadline(); + if (ec || len != sizeof(status_frame_)) { + // The body path may have already finalized a failure and + // closed the socket; finalize() is idempotent (moved-from + // callbacks are null-checked). + if (ec != asio::error::operation_aborted) { + LOG(ERROR) + << "ClientSession: failed to read WRITE " + "ack frame. Error: " + << ec.message() << " (value: " << ec.value() << ")"; + } + abortWrite(); + return; + } + uint64_t frame = le64toh(status_frame_); + if (!statusFrameValid(frame)) { + LOG(ERROR) << "ClientSession: malformed WRITE ack frame " + "(peer likely speaks the legacy protocol); " + "dropping connection"; + abortWrite(); + return; + } + if (frame != kStatusOk) { + LOG(ERROR) << "ClientSession: WRITE rejected by server, " + "status " + << (frame & 0xFFFFFFFFull); + abortWrite(); + return; + } + if (!write_body_done_) { + // The server's ack can legitimately overtake the final + // local write-completion handler (both become ready + // together for small writes; the io thread may run this + // handler first). Record it; the body path finalizes. + write_acked_ok_ = true; + return; + } + finalize(TransferStatusEnum::COMPLETED, true); + }); + } + + void readBody() { + auto self(shared_from_this()); + uint64_t size = le64toh(header_.size); + char* addr = local_buffer_; + + size_t buffer_size = + std::min(getChunkSize(), size - total_transferred_bytes_); + if (buffer_size == 0) { + finalize(TransferStatusEnum::COMPLETED, true); + return; + } + + char* dram_buffer = addr + total_transferred_bytes_; + int cuda_device = -1; + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + cuda_device = getCudaDeviceId(addr); + if (cuda_device >= 0) { + dram_buffer = new char[buffer_size]; + } +#endif + + asio::async_read( + *socket_, asio::buffer(dram_buffer, buffer_size), + [this, addr, dram_buffer, cuda_device, self]( + const asio::error_code& ec, std::size_t transferred_bytes) { + if (ec) { + LOG(ERROR) + << "ClientSession::readBody failed. " + << "Attempt to read data " << static_cast(addr) + << " using buffer " << static_cast(dram_buffer) + << ". Error: " << ec.message() + << " (value: " << ec.value() << ")"; +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + if (cuda_device >= 0) delete[] dram_buffer; +#endif + finalize(TransferStatusEnum::FAILED, false); + return; + } + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + if (cuda_device >= 0) { + cudaSetDevice(cuda_device); +#ifdef USE_MACA + cudaError_t cuda_status = + copyTcpCudaMemory(addr + total_transferred_bytes_, + dram_buffer, transferred_bytes); +#else + cudaError_t cuda_status = + cudaMemcpy(addr + total_transferred_bytes_, dram_buffer, + transferred_bytes, cudaMemcpyDefault); +#endif + if (cuda_status != cudaSuccess) { + LOG(ERROR) + << "ClientSession::readBody failed to copy to CUDA " + "memory. " + << "Error: " << cudaGetErrorString(cuda_status); + delete[] dram_buffer; + finalize(TransferStatusEnum::FAILED, false); + return; + } + delete[] dram_buffer; + } +#endif + total_transferred_bytes_ += transferred_bytes; + readBody(); + }); + } + + void writeBody() { + auto self(shared_from_this()); + uint64_t size = le64toh(header_.size); + char* addr = local_buffer_; + + size_t buffer_size = + std::min(getChunkSize(), size - total_transferred_bytes_); + if (buffer_size == 0) { + if (v2_) { + if (write_abort_requested_) { + finalize(TransferStatusEnum::FAILED, false); + return; + } + // Completion comes from the server's acknowledgment, whose + // read is already in flight (armed in writeHeader) and may + // have finished first. + write_body_done_ = true; + if (write_acked_ok_) { + finalize(TransferStatusEnum::COMPLETED, true); + } else { + // From here a well-behaved server owes at most one + // chunk's apply plus the frame; a legacy peer behind a + // stale descriptor may owe nothing, ever. + armStatusDeadline(); + } + } else { + // v1: no acknowledgment exists in the protocol; this only + // means the payload left the initiator (#2086). + finalize(TransferStatusEnum::COMPLETED, true); + } + return; + } + + char* dram_buffer = addr + total_transferred_bytes_; + int cuda_device = -1; + +#if defined(USE_CUDA) || defined(USE_MUSA) || defined(USE_HIP) || \ + defined(USE_MLU) || defined(USE_MACA) || defined(USE_HYGON) || \ + defined(USE_COREX) + cuda_device = getCudaDeviceId(addr); + if (cuda_device >= 0) { + dram_buffer = new char[buffer_size]; + cudaSetDevice(cuda_device); +#ifdef USE_MACA + cudaError_t cuda_status = copyTcpCudaMemory( + dram_buffer, addr + total_transferred_bytes_, buffer_size); +#else + cudaError_t cuda_status = + cudaMemcpy(dram_buffer, addr + total_transferred_bytes_, + buffer_size, cudaMemcpyDefault); +#endif + if (cuda_status != cudaSuccess) { + LOG(ERROR) << "ClientSession::writeBody failed to copy from " + "CUDA memory. " + << "Error: " << cudaGetErrorString(cuda_status); + delete[] dram_buffer; + abortWrite(); + return; + } + } +#endif + + write_body_in_flight_ = true; + asio::async_write( + *socket_, asio::buffer(dram_buffer, buffer_size), + [this, addr, dram_buffer, cuda_device, self]( + const asio::error_code& ec, std::size_t transferred_bytes) { + write_body_in_flight_ = false; + if (cuda_device >= 0) { + delete[] dram_buffer; + } + if (ec) { + LOG(ERROR) + << "ClientSession::writeBody failed. " + << "Attempt to write data " << static_cast(addr) + << " using buffer " << static_cast(dram_buffer) + << ". Error: " << ec.message() + << " (value: " << ec.value() << ")"; + abortWrite(); + return; + } + if (write_abort_requested_) { + // The early ack path closed the socket while this + // operation still owned the caller's source buffer. It is + // safe to publish failure now that the handler has run. + finalize(TransferStatusEnum::FAILED, false); + return; + } + total_transferred_bytes_ += transferred_bytes; + writeBody(); + }); + } +}; + +struct TcpContext { + TcpContext(short port, ValidateAddrFn validate_addr) + : acceptor(io_context), validate_addr_(std::move(validate_addr)) { + std::error_code ec; + asio::ip::tcp::endpoint endpoint(asio::ip::tcp::v6(), port); + + acceptor.open(endpoint.protocol(), ec); + if (!ec) { + acceptor.set_option(asio::ip::v6_only(false), ec); + if (!ec) { + acceptor.set_option( + asio::ip::tcp::acceptor::reuse_address(true)); + acceptor.bind(endpoint, ec); + if (!ec) { + acceptor.listen(); + return; + } + } + acceptor.close(); + } + LOG(ERROR) << "Failed to set up IPv6 dual-stack listener: " + << ec.message() << " (error code: " << ec.value() << ")"; + asio::ip::tcp::endpoint endpoint_v4(asio::ip::tcp::v4(), port); + acceptor.open(endpoint_v4.protocol()); + acceptor.set_option(asio::ip::tcp::acceptor::reuse_address(true)); + acceptor.bind(endpoint_v4); + acceptor.listen(); + } + + void doAccept() { + acceptor.async_accept([this](asio::error_code ec, tcpsocket socket) { + if (!ec) { + asio::error_code nodelay_ec; + socket.set_option(asio::ip::tcp::no_delay(true), nodelay_ec); + auto socket_ptr = + std::make_shared(std::move(socket)); + auto session = + std::make_shared(socket_ptr, validate_addr_); + session->start(); + } + doAccept(); + }); + } + + asio::io_context io_context; + asio::ip::tcp::acceptor acceptor; + ValidateAddrFn validate_addr_; +}; diff --git a/mooncake-transfer-engine/tests/CMakeLists.txt b/mooncake-transfer-engine/tests/CMakeLists.txt index aa00a168e1..6935fd7e96 100644 --- a/mooncake-transfer-engine/tests/CMakeLists.txt +++ b/mooncake-transfer-engine/tests/CMakeLists.txt @@ -113,6 +113,8 @@ if(USE_TCP) add_executable(tcp_write_visibility_test ${WORKSPACE}/tcp_write_visibility_test.cpp) + target_compile_definitions(tcp_write_visibility_test + PRIVATE MOONCAKE_TCP_TRANSPORT_TEST_HOOKS) target_link_libraries(tcp_write_visibility_test PUBLIC transfer_engine gtest gtest_main) add_test(NAME tcp_write_visibility_test COMMAND tcp_write_visibility_test) diff --git a/mooncake-transfer-engine/tests/tcp_write_visibility_test.cpp b/mooncake-transfer-engine/tests/tcp_write_visibility_test.cpp index 6bc9d23408..4af53322fd 100644 --- a/mooncake-transfer-engine/tests/tcp_write_visibility_test.cpp +++ b/mooncake-transfer-engine/tests/tcp_write_visibility_test.cpp @@ -39,7 +39,9 @@ #include #include #include +#include #include +#include #include #include #include @@ -47,6 +49,24 @@ #include "transfer_engine.h" #include "transport/transport.h" +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +namespace mooncake { +void tcpTransportSetLaneConnectHandlerHookForTest( + void (*hook)() noexcept) noexcept; +void tcpTransportSetLaneConnectFailureInjectionHookForTest( + bool (*hook)(size_t) noexcept) noexcept; +void tcpTransportSetLaneObserverHookForTest( + void (*hook)(int, size_t, uint64_t, size_t, bool) noexcept) noexcept; +void tcpTransportSetLaneRetryHandlerHookForTest( + void (*hook)() noexcept) noexcept; +void tcpTransportSetLaneAdmissionHandlerHookForTest( + void (*hook)() noexcept) noexcept; +void tcpTransportSetLaneFailureReasonHookForTest( + void (*hook)(int) noexcept) noexcept; +bool tcpTransportLaneTypesAreMoveOnlyForTest() noexcept; +} // namespace mooncake +#endif + using namespace mooncake; namespace { @@ -57,6 +77,333 @@ constexpr size_t kRegionAlign = 4 * 1024 * 1024; constexpr int kIterations = 400; constexpr int kNoiseThreads = 3; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +enum LaneTestEvent { + kLaneQueueAdmitted = 1, + kLaneQueueRejected = 2, + kLaneConnecting = 3, + kLaneBusy = 4, + kLaneTerminal = 5, + kLaneShutdownClean = 6, + kLaneLateHandler = 7, + kLaneRetryArmed = 8, + kLaneRetryFired = 9, + kLaneRetryLate = 10, + kLaneCooldownStarted = 11, + kLaneAdmissionPending = 12, + kLaneAdmissionPromoted = 13, + kLaneAdmissionTimerArmed = 14, + kLaneAdmissionTimerFired = 15, + kLaneAdmissionTimerLate = 16, + kLaneAdmissionHardRejected = 17, +}; + +enum WorkFailureReasonForTest { + kWorkQueueFull = 0, + kWorkQueueTimeout = 1, + kWorkRuntimeUnavailable = 2, + kWorkConnectFailed = 3, + kWorkSessionFailed = 4, + kWorkShutdown = 5, +}; + +std::atomic lane_connect_handler_entered{false}; +std::atomic lane_connect_injection_call_count{0}; +std::atomic release_lane_connect_handler{false}; +std::atomic hold_lane_connect_handler{false}; +std::atomic hold_lane_connect_after_busy{false}; +std::atomic connecting_lane_had_current{false}; +std::atomic retry_handler_entered{false}; +std::atomic release_retry_handler{false}; +std::atomic hold_retry_handler{false}; +std::atomic retry_armed_observer_entered{false}; +std::atomic release_retry_armed_observer{false}; +std::atomic hold_retry_armed_observer{false}; +std::atomic admission_handler_entered{false}; +std::atomic release_admission_handler{false}; +std::atomic hold_admission_handler{false}; +std::atomic maximum_observed_queue_depth{0}; +std::atomic maximum_observed_pending_depth{0}; +std::atomic maximum_observed_socket_count{0}; +std::atomic queue_rejection_count{0}; +std::atomic lane_connecting_count{0}; +std::atomic lane_busy_count{0}; +std::atomic lane_terminal_count{0}; +std::atomic lane_shutdown_clean_count{0}; +std::atomic late_lane_handler_count{0}; +std::atomic retry_armed_count{0}; +std::atomic retry_fired_count{0}; +std::atomic retry_late_count{0}; +std::atomic cooldown_started_count{0}; +std::atomic admission_pending_count{0}; +std::atomic admission_promoted_count{0}; +std::atomic admission_timer_armed_count{0}; +std::atomic admission_timer_fired_count{0}; +std::atomic admission_timer_late_count{0}; +std::atomic admission_hard_rejection_count{0}; +std::atomic queue_full_failure_count{0}; +std::atomic queue_timeout_failure_count{0}; +std::atomic connect_failure_count{0}; +std::atomic runtime_unavailable_failure_count{0}; +std::atomic session_failure_count{0}; +std::atomic shutdown_failure_count{0}; + +template +bool waitForPredicate(Predicate&& predicate, + std::chrono::duration timeout) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (predicate()) return true; + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return predicate(); +} + +template +void updateMaximum(std::atomic& value, T candidate) noexcept { + T previous = value.load(std::memory_order_relaxed); + while (previous < candidate && + !value.compare_exchange_weak(previous, candidate, + std::memory_order_relaxed)) { + } +} + +void resetLaneTestState() noexcept { + lane_connect_handler_entered.store(false, std::memory_order_release); + lane_connect_injection_call_count.store(0, std::memory_order_release); + release_lane_connect_handler.store(false, std::memory_order_release); + hold_lane_connect_handler.store(false, std::memory_order_release); + hold_lane_connect_after_busy.store(false, std::memory_order_release); + connecting_lane_had_current.store(false, std::memory_order_release); + retry_handler_entered.store(false, std::memory_order_release); + release_retry_handler.store(false, std::memory_order_release); + hold_retry_handler.store(false, std::memory_order_release); + retry_armed_observer_entered.store(false, std::memory_order_release); + release_retry_armed_observer.store(false, std::memory_order_release); + hold_retry_armed_observer.store(false, std::memory_order_release); + admission_handler_entered.store(false, std::memory_order_release); + release_admission_handler.store(false, std::memory_order_release); + hold_admission_handler.store(false, std::memory_order_release); + maximum_observed_queue_depth.store(0, std::memory_order_release); + maximum_observed_pending_depth.store(0, std::memory_order_release); + maximum_observed_socket_count.store(0, std::memory_order_release); + queue_rejection_count.store(0, std::memory_order_release); + lane_connecting_count.store(0, std::memory_order_release); + lane_busy_count.store(0, std::memory_order_release); + lane_terminal_count.store(0, std::memory_order_release); + lane_shutdown_clean_count.store(0, std::memory_order_release); + late_lane_handler_count.store(0, std::memory_order_release); + retry_armed_count.store(0, std::memory_order_release); + retry_fired_count.store(0, std::memory_order_release); + retry_late_count.store(0, std::memory_order_release); + cooldown_started_count.store(0, std::memory_order_release); + admission_pending_count.store(0, std::memory_order_release); + admission_promoted_count.store(0, std::memory_order_release); + admission_timer_armed_count.store(0, std::memory_order_release); + admission_timer_fired_count.store(0, std::memory_order_release); + admission_timer_late_count.store(0, std::memory_order_release); + admission_hard_rejection_count.store(0, std::memory_order_release); + queue_full_failure_count.store(0, std::memory_order_release); + queue_timeout_failure_count.store(0, std::memory_order_release); + connect_failure_count.store(0, std::memory_order_release); + runtime_unavailable_failure_count.store(0, std::memory_order_release); + session_failure_count.store(0, std::memory_order_release); + shutdown_failure_count.store(0, std::memory_order_release); +} + +bool failSecondLaneConnectAttemptOnce(size_t) noexcept { + const int attempt = lane_connect_injection_call_count.fetch_add(1) + 1; + return attempt == 2; +} + +bool failLaneOneAndTwoConnectAttempt(size_t lane_id) noexcept { + lane_connect_injection_call_count.fetch_add(1, std::memory_order_relaxed); + return lane_id == 1 || lane_id == 2; +} + +void releaseLaneConnectHandler() noexcept { + release_lane_connect_handler.store(true, std::memory_order_release); +} + +void blockLaneConnectHandler() noexcept { + if (!hold_lane_connect_handler.load(std::memory_order_acquire)) return; + lane_connect_handler_entered.store(true, std::memory_order_release); + while (!release_lane_connect_handler.load(std::memory_order_acquire)) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } +} + +void releaseRetryHandler() noexcept { + release_retry_handler.store(true, std::memory_order_release); +} + +void blockRetryHandler() noexcept { + if (!hold_retry_handler.load(std::memory_order_acquire)) return; + retry_handler_entered.store(true, std::memory_order_release); + while (!release_retry_handler.load(std::memory_order_acquire)) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } +} + +void releaseRetryArmedObserver() noexcept { + release_retry_armed_observer.store(true, std::memory_order_release); +} + +void releaseAdmissionHandler() noexcept { + release_admission_handler.store(true, std::memory_order_release); +} + +void blockAdmissionHandler() noexcept { + if (!hold_admission_handler.load(std::memory_order_acquire)) return; + admission_handler_entered.store(true, std::memory_order_release); + while (!release_admission_handler.load(std::memory_order_acquire)) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } +} + +void observeLaneState(int event, size_t queue_depth, uint64_t, + size_t active_sockets, bool lane_has_current) noexcept { + if (event < kLaneAdmissionPending || event > kLaneAdmissionHardRejected) + updateMaximum(maximum_observed_queue_depth, queue_depth); + updateMaximum(maximum_observed_socket_count, active_sockets); + if (event == kLaneConnecting && lane_has_current) + connecting_lane_had_current.store(true, std::memory_order_release); + if (event == kLaneConnecting) + lane_connecting_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneBusy) { + lane_busy_count.fetch_add(1, std::memory_order_relaxed); + if (hold_lane_connect_after_busy.load(std::memory_order_acquire)) + hold_lane_connect_handler.store(true, std::memory_order_release); + } + if (event == kLaneQueueRejected) + queue_rejection_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneTerminal) + lane_terminal_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneShutdownClean) + lane_shutdown_clean_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneLateHandler) + late_lane_handler_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneRetryArmed) { + retry_armed_count.fetch_add(1, std::memory_order_relaxed); + if (hold_retry_armed_observer.load(std::memory_order_acquire)) { + retry_armed_observer_entered.store(true, std::memory_order_release); + while ( + !release_retry_armed_observer.load(std::memory_order_acquire)) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + } + } + if (event == kLaneRetryFired) + retry_fired_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneRetryLate) + retry_late_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneCooldownStarted) + cooldown_started_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneAdmissionPending) { + admission_pending_count.fetch_add(1, std::memory_order_relaxed); + updateMaximum(maximum_observed_pending_depth, queue_depth); + } + if (event == kLaneAdmissionPromoted) { + admission_promoted_count.fetch_add(1, std::memory_order_relaxed); + updateMaximum(maximum_observed_pending_depth, queue_depth); + } + if (event == kLaneAdmissionTimerArmed) + admission_timer_armed_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneAdmissionTimerFired) + admission_timer_fired_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneAdmissionTimerLate) + admission_timer_late_count.fetch_add(1, std::memory_order_relaxed); + if (event == kLaneAdmissionHardRejected) + admission_hard_rejection_count.fetch_add(1, std::memory_order_relaxed); +} + +void observeWorkFailureReason(int reason) noexcept { + if (reason == kWorkQueueFull) + queue_full_failure_count.fetch_add(1, std::memory_order_relaxed); + if (reason == kWorkQueueTimeout) + queue_timeout_failure_count.fetch_add(1, std::memory_order_relaxed); + if (reason == kWorkConnectFailed) + connect_failure_count.fetch_add(1, std::memory_order_relaxed); + if (reason == kWorkRuntimeUnavailable) + runtime_unavailable_failure_count.fetch_add(1, + std::memory_order_relaxed); + if (reason == kWorkSessionFailed) + session_failure_count.fetch_add(1, std::memory_order_relaxed); + if (reason == kWorkShutdown) + shutdown_failure_count.fetch_add(1, std::memory_order_relaxed); +} + +class ScopedLaneHooks { + public: + explicit ScopedLaneHooks(bool block_first_connect_handler = false, + bool block_after_busy = false, + bool block_retry = false, + bool block_retry_armed = false, + bool block_admission = false, + bool fail_second_connect = false) { + resetLaneTestState(); + tcpTransportSetLaneObserverHookForTest(observeLaneState); + tcpTransportSetLaneFailureReasonHookForTest(observeWorkFailureReason); + if (fail_second_connect) { + tcpTransportSetLaneConnectFailureInjectionHookForTest( + failSecondLaneConnectAttemptOnce); + } + if (block_first_connect_handler || block_after_busy) { + hold_lane_connect_handler.store(block_first_connect_handler, + std::memory_order_release); + hold_lane_connect_after_busy.store(block_after_busy, + std::memory_order_release); + tcpTransportSetLaneConnectHandlerHookForTest( + blockLaneConnectHandler); + } + if (block_retry) { + hold_retry_handler.store(true, std::memory_order_release); + tcpTransportSetLaneRetryHandlerHookForTest(blockRetryHandler); + } + hold_retry_armed_observer.store(block_retry_armed, + std::memory_order_release); + if (block_admission) { + hold_admission_handler.store(true, std::memory_order_release); + tcpTransportSetLaneAdmissionHandlerHookForTest( + blockAdmissionHandler); + } + } + + ~ScopedLaneHooks() { reset(); } + + void reset() noexcept { + if (!active_) return; + releaseLaneConnectHandler(); + releaseRetryHandler(); + releaseRetryArmedObserver(); + releaseAdmissionHandler(); + tcpTransportSetLaneConnectHandlerHookForTest(nullptr); + tcpTransportSetLaneConnectFailureInjectionHookForTest(nullptr); + tcpTransportSetLaneRetryHandlerHookForTest(nullptr); + tcpTransportSetLaneAdmissionHandlerHookForTest(nullptr); + tcpTransportSetLaneObserverHookForTest(nullptr); + tcpTransportSetLaneFailureReasonHookForTest(nullptr); + active_ = false; + } + + private: + bool active_ = true; +}; +#endif + +void reclaimBatchDescAfterEngineShutdownForTest(Transport::BatchID batch_id) { +#ifndef CONFIG_USE_BATCH_DESC_SET + // These shutdown tests deliberately destroy TransferEngine with work still + // outstanding, so normal freeBatchID cannot be called. The caller must + // wait for engine destruction to return (and thus for its worker to join) + // before reclaiming a descriptor that no callback can touch anymore. + delete &Transport::toBatchDesc(batch_id); +#else + // The batch descriptor registry owns (and may already reclaim) this object. + (void)batch_id; +#endif +} + class ScopedEnvVar { public: ScopedEnvVar(const char* name, const char* value) : name_(name) { @@ -209,6 +556,350 @@ class LegacyReadServer { std::atomic saw_flagged_write_{false}; }; +// Accepts TCP connections but deliberately never reads request headers or +// bodies. Large writes therefore remain in progress until closePeer() drops +// the accepted sockets. +class HoldingWriteServer { + public: + HoldingWriteServer() { + listen_fd_ = socket(AF_INET, SOCK_STREAM, 0); + if (listen_fd_ < 0) return; + int one = 1; + if (setsockopt(listen_fd_, SOL_SOCKET, SO_REUSEADDR, &one, + sizeof(one)) != 0) + return; + + sockaddr_in addr{}; + addr.sin_family = AF_INET; + addr.sin_port = 0; + addr.sin_addr.s_addr = htonl(INADDR_ANY); + if (bind(listen_fd_, reinterpret_cast(&addr), + sizeof(addr)) != 0) + return; + if (listen(listen_fd_, 64) != 0) return; + + socklen_t len = sizeof(addr); + if (getsockname(listen_fd_, reinterpret_cast(&addr), &len) != + 0) + return; + port_ = ntohs(addr.sin_port); + ok_ = true; + thread_ = std::thread([this] { acceptLoop(); }); + } + + ~HoldingWriteServer() { + closePeer(); + if (listen_fd_ >= 0) close(listen_fd_); + } + + bool ok() const { return ok_; } + uint16_t port() const { return port_; } + int acceptedCount() const { return accepted_count_.load(); } + int maxAcceptedCount() const { return max_accepted_count_.load(); } + + bool waitForAccepted(int count, std::chrono::seconds timeout) const { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (acceptedCount() < count && + std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return acceptedCount() >= count; + } + + int activeAcceptedCount() const { + std::lock_guard lock(accepted_mutex_); + return static_cast(accepted_fds_.size()); + } + + bool closeOneAccepted() { + int fd = -1; + { + std::lock_guard lock(accepted_mutex_); + if (accepted_fds_.empty()) return false; + fd = accepted_fds_.front(); + accepted_fds_.erase(accepted_fds_.begin()); + } + (void)shutdown(fd, SHUT_RDWR); + close(fd); + return true; + } + + void closePeer() { + if (closed_.exchange(true)) return; + if (listen_fd_ >= 0) (void)shutdown(listen_fd_, SHUT_RDWR); + if (thread_.joinable()) thread_.join(); + + std::vector accepted; + { + std::lock_guard lock(accepted_mutex_); + accepted.swap(accepted_fds_); + } + for (int fd : accepted) { + (void)shutdown(fd, SHUT_RDWR); + close(fd); + } + } + + private: + void acceptLoop() { + while (!closed_.load()) { + int fd = accept(listen_fd_, nullptr, nullptr); + if (fd < 0) break; + if (closed_.load()) { + close(fd); + break; + } + + { + std::lock_guard lock(accepted_mutex_); + accepted_fds_.push_back(fd); + } + const int count = accepted_count_.fetch_add(1) + 1; + int previous_max = max_accepted_count_.load(); + while (previous_max < count && + !max_accepted_count_.compare_exchange_weak(previous_max, + count)) { + } + } + } + + int listen_fd_ = -1; + uint16_t port_ = 0; + bool ok_ = false; + std::thread thread_; + mutable std::mutex accepted_mutex_; + std::vector accepted_fds_; + std::atomic closed_{false}; + std::atomic accepted_count_{0}; + std::atomic max_accepted_count_{0}; +}; + +// Keeps a local TCP port bound but deliberately never listens. Linux rejects +// connect attempts deterministically, without blackhole-address timing. +class UnavailableTcpPeer { + public: + UnavailableTcpPeer() { + fd_ = socket(AF_INET, SOCK_STREAM, 0); + if (fd_ < 0) return; + + sockaddr_in addr{}; + addr.sin_family = AF_INET; + addr.sin_port = 0; + addr.sin_addr.s_addr = htonl(INADDR_ANY); + if (bind(fd_, reinterpret_cast(&addr), sizeof(addr)) != 0) + return; + + socklen_t len = sizeof(addr); + if (getsockname(fd_, reinterpret_cast(&addr), &len) != 0) + return; + port_ = ntohs(addr.sin_port); + ok_ = true; + } + + ~UnavailableTcpPeer() { + if (fd_ >= 0) close(fd_); + } + + bool ok() const { return ok_; } + uint16_t port() const { return port_; } + + private: + int fd_ = -1; + uint16_t port_ = 0; + bool ok_ = false; +}; + +// Minimal v2 WRITE server that keeps accepted connections open and processes +// multiple request/ack exchanges on each one. +class ReusingWriteServer { + public: + explicit ReusingWriteServer(bool hold_first_ack = false) + : hold_first_ack_(hold_first_ack) { + listen_fd_ = socket(AF_INET, SOCK_STREAM, 0); + if (listen_fd_ < 0) return; + int one = 1; + if (setsockopt(listen_fd_, SOL_SOCKET, SO_REUSEADDR, &one, + sizeof(one)) != 0) + return; + + sockaddr_in addr{}; + addr.sin_family = AF_INET; + addr.sin_port = 0; + addr.sin_addr.s_addr = htonl(INADDR_ANY); + if (bind(listen_fd_, reinterpret_cast(&addr), + sizeof(addr)) != 0) + return; + if (listen(listen_fd_, 16) != 0) return; + + socklen_t len = sizeof(addr); + if (getsockname(listen_fd_, reinterpret_cast(&addr), &len) != + 0) + return; + port_ = ntohs(addr.sin_port); + ok_ = true; + accept_thread_ = std::thread([this] { acceptLoop(); }); + } + + ~ReusingWriteServer() { stop(); } + + bool ok() const { return ok_; } + uint16_t port() const { return port_; } + int acceptedCount() const { return accepted_count_.load(); } + + bool waitForRequests(int count, std::chrono::seconds timeout) const { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (request_count_.load() < count && + std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return request_count_.load() >= count; + } + + bool waitForFirstRequest(std::chrono::seconds timeout) const { + return waitForPredicate( + [this] { + return first_request_received_.load(std::memory_order_acquire); + }, + timeout); + } + + void releaseFirstAck() noexcept { + release_first_ack_.store(true, std::memory_order_release); + } + + bool waitForFirstAckSent(std::chrono::seconds timeout) const { + return waitForPredicate( + [this] { return first_ack_sent_.load(std::memory_order_acquire); }, + timeout); + } + + std::vector requestAddresses() const { + std::lock_guard lock(request_mutex_); + return request_addresses_; + } + + private: + static bool recvExact(int fd, void* buffer, size_t size) { + char* out = static_cast(buffer); + while (size) { + ssize_t n = recv(fd, out, size, 0); + if (n <= 0) return false; + out += n; + size -= static_cast(n); + } + return true; + } + + static bool sendExact(int fd, const void* buffer, size_t size) { + const char* in = static_cast(buffer); + while (size) { + ssize_t n = send(fd, in, size, MSG_NOSIGNAL); + if (n <= 0) return false; + in += n; + size -= static_cast(n); + } + return true; + } + + void acceptLoop() { + while (!stopped_.load()) { + int fd = accept(listen_fd_, nullptr, nullptr); + if (fd < 0) break; + if (stopped_.load()) { + close(fd); + break; + } + { + std::lock_guard lock(connection_mutex_); + accepted_fds_.push_back(fd); + connection_threads_.emplace_back( + [this, fd] { serveConnection(fd); }); + } + accepted_count_.fetch_add(1); + } + } + + void serveConnection(int fd) { + std::vector payload(64 * 1024); + bool keep_serving = true; + while (!stopped_.load()) { + TestSessionHeader header{}; + if (!recvExact(fd, &header, sizeof(header))) break; + { + std::lock_guard lock(request_mutex_); + request_addresses_.push_back(le64toh(header.addr)); + } + uint64_t remaining = le64toh(header.size); + while (remaining) { + const size_t chunk = + std::min(payload.size(), remaining); + if (!recvExact(fd, payload.data(), chunk)) { + keep_serving = false; + break; + } + remaining -= chunk; + } + if (!keep_serving) break; + + const int request_index = + requests_received_.fetch_add(1, std::memory_order_acq_rel); + if (hold_first_ack_ && request_index == 0) { + first_request_received_.store(true, std::memory_order_release); + while (!release_first_ack_.load(std::memory_order_acquire) && + !stopped_.load(std::memory_order_acquire)) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + } + + constexpr uint64_t kStatusOk = 0x4D435456ull << 32; + const uint64_t status = htole64(kStatusOk); + if (!sendExact(fd, &status, sizeof(status))) break; + if (hold_first_ack_ && request_index == 0) + first_ack_sent_.store(true, std::memory_order_release); + request_count_.fetch_add(1); + } + { + std::lock_guard lock(connection_mutex_); + auto it = std::find(accepted_fds_.begin(), accepted_fds_.end(), fd); + if (it != accepted_fds_.end()) accepted_fds_.erase(it); + } + close(fd); + } + + void stop() { + if (stopped_.exchange(true)) return; + releaseFirstAck(); + if (listen_fd_ >= 0) (void)shutdown(listen_fd_, SHUT_RDWR); + if (accept_thread_.joinable()) accept_thread_.join(); + + { + std::lock_guard lock(connection_mutex_); + for (int fd : accepted_fds_) (void)shutdown(fd, SHUT_RDWR); + } + for (auto& thread : connection_threads_) + if (thread.joinable()) thread.join(); + if (listen_fd_ >= 0) close(listen_fd_); + } + + int listen_fd_ = -1; + uint16_t port_ = 0; + bool ok_ = false; + std::thread accept_thread_; + std::mutex connection_mutex_; + mutable std::mutex request_mutex_; + std::vector accepted_fds_; + std::vector connection_threads_; + std::vector request_addresses_; + std::atomic stopped_{false}; + std::atomic accepted_count_{0}; + std::atomic request_count_{0}; + const bool hold_first_ack_; + std::atomic requests_received_{0}; + std::atomic first_request_received_{false}; + std::atomic release_first_ack_{false}; + std::atomic first_ack_sent_{false}; +}; + struct EngineHandle { std::unique_ptr engine; void* pool = nullptr; @@ -254,6 +945,26 @@ struct EngineHandle { } }; +void pointTcpSegmentAt(EngineHandle& handle, uint16_t port) { + auto desc = + handle.engine->getMetadata()->getSegmentDescByID(handle.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = port; + desc->tcp_proto_version = 2; +} + +TransferRequest makeWriteRequest(const EngineHandle& handle, size_t length, + uint64_t target_offset = 0) { + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = length; + request.source = handle.pool; + request.target_id = handle.segment_id; + request.target_offset = + target_offset == 0 ? handle.remote_base : target_offset; + return request; +} + // Submit one request and poll until terminal state; returns final status. TransferStatusEnum runOne(TransferEngine* engine, TransferRequest entry) { auto batch_id = engine->allocateBatchID(1); @@ -274,6 +985,75 @@ TransferStatusEnum runOne(TransferEngine* engine, TransferRequest entry) { return status.s; } +bool waitForBatchTerminal(TransferEngine* engine, Transport::BatchID batch_id, + size_t task_count, std::chrono::seconds timeout) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + bool all_terminal = true; + for (size_t task_id = 0; task_id < task_count; ++task_id) { + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + if (!engine->getTransferStatus(batch_id, task_id, status).ok() || + (status.s != TransferStatusEnum::COMPLETED && + status.s != TransferStatusEnum::FAILED)) { + all_terminal = false; + } + } + if (all_terminal) return true; + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + + for (size_t task_id = 0; task_id < task_count; ++task_id) { + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + if (!engine->getTransferStatus(batch_id, task_id, status).ok() || + (status.s != TransferStatusEnum::COMPLETED && + status.s != TransferStatusEnum::FAILED)) { + return false; + } + } + return true; +} + +void expectEverySliceCompletedExactlyOnceAfterShutdown( + Transport::BatchID batch_id) { + const auto& batch = Transport::toBatchDesc(batch_id); + for (size_t task_id = 0; task_id < batch.task_list.size(); ++task_id) { + const auto& task = batch.task_list[task_id]; + const uint64_t success = + __atomic_load_n(&task.success_slice_count, __ATOMIC_RELAXED); + const uint64_t failed = + __atomic_load_n(&task.failed_slice_count, __ATOMIC_RELAXED); + const uint64_t slices = + __atomic_load_n(&task.slice_count, __ATOMIC_RELAXED); + EXPECT_EQ(success + failed, slices) << "task " << task_id; + for (const auto* slice : task.slice_list) { + EXPECT_TRUE(slice->status == Transport::Slice::SUCCESS || + slice->status == Transport::Slice::FAILED) + << "task " << task_id; + } + } +} + +void expectEverySliceSucceededExactlyOnceAfterShutdown( + Transport::BatchID batch_id) { + const auto& batch = Transport::toBatchDesc(batch_id); + for (size_t task_id = 0; task_id < batch.task_list.size(); ++task_id) { + const auto& task = batch.task_list[task_id]; + const uint64_t success = + __atomic_load_n(&task.success_slice_count, __ATOMIC_RELAXED); + const uint64_t failed = + __atomic_load_n(&task.failed_slice_count, __ATOMIC_RELAXED); + const uint64_t slices = + __atomic_load_n(&task.slice_count, __ATOMIC_RELAXED); + EXPECT_EQ(success, slices) << "task " << task_id; + EXPECT_EQ(failed, 0u) << "task " << task_id; + for (const auto* slice : task.slice_list) + EXPECT_EQ(slice->status, Transport::Slice::SUCCESS) + << "task " << task_id; + } +} + } // namespace TEST(TcpWriteVisibilityTest, CompletedWriteIsVisibleToSubsequentRead) { @@ -486,16 +1266,26 @@ TEST(TcpWriteVisibilityTest, V2ReadRoundTripAndRejectedRead) { EXPECT_EQ(runOne(h.engine.get(), r), TransferStatusEnum::COMPLETED); } -// Mixed-version quadrant: a legacy (v1) initiator against the v2 server — -// selected via the MC_TCP_PROTO=1 escape hatch — still transfers data -// correctly (with the old weaker completion semantics). -TEST(TcpWriteVisibilityTest, LegacyInitiatorInteropWithV2Server) { +// Mixed-version quadrant: a pooled legacy (v1) initiator against the current +// server still transfers data correctly over repeated exchanges on one fixed +// lane (with the old weaker WRITE completion semantics). +TEST(TcpWriteVisibilityTest, + PooledLegacyInitiatorInteroperatesWithCurrentServer) { const char* env = std::getenv("MC_METADATA_SERVER"); std::string metadata_server = env ? env : "P2PHANDSHAKE"; // Set the process environment before the engine starts any threads, and // restore it only after those threads have stopped. POSIX does not require // setenv()/unsetenv() to synchronize with concurrent getenv() calls. + ScopedEnvVar pooling("MC_TCP_ENABLE_CONNECTION_POOL", "1"); + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "2"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "2"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); ScopedEnvVar legacy_proto("MC_TCP_PROTO", "1"); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + ScopedLaneHooks hooks; +#endif { EngineHandle h; h.init(metadata_server, "127.0.0.2:17903", 16ull << 20); @@ -531,7 +1321,15 @@ TEST(TcpWriteVisibilityTest, LegacyInitiatorInteropWithV2Server) { r.target_offset = h.remote_base; EXPECT_EQ(runOne(h.engine.get(), r), TransferStatusEnum::COMPLETED); EXPECT_EQ(memcmp(dst, src, kSmallLength), 0); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 1u); +#endif } +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + hooks.reset(); +#endif } // A cached v2 descriptor can briefly outlive a server downgrade/restart. A @@ -636,3 +1434,1447 @@ TEST(TcpWriteVisibilityTest, StaleV2DescriptorShortRequestFailsWithinDeadline) { legacy_server.join(); } } + +TEST(TcpWriteVisibilityTest, PerPeerLaneAndQueueBoundsHoldUnderLoad) { + constexpr int kRounds = 3; + constexpr int kRequestCount = 32; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "2"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "64"); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + for (int round = 0; round < kRounds; ++round) { +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + ScopedLaneHooks hooks; +#endif + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()) << "round " << round; + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:" + std::to_string(17907 + round), + kLength); + ASSERT_TRUE(h.ok) << "engine/segment setup failed in round " << round; + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = fake_peer.port(); + desc->tcp_proto_version = 2; + + memset(h.pool, 0x4D + round, kLength); + std::vector requests(kRequestCount); + for (auto& request : requests) { + request.opcode = TransferRequest::WRITE; + request.length = kLength; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + } + + auto batch_id = h.engine->allocateBatchID(kRequestCount); + Status submission = h.engine->submitTransfer(batch_id, requests); + ASSERT_TRUE(submission.ok()) << "round " << round; + + EXPECT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))) + << "round " << round << " accepted only " + << fake_peer.acceptedCount() << " connections"; + // Give any incorrectly uncapped attempts time to reach accept(). + std::this_thread::sleep_for(std::chrono::milliseconds(250)); + EXPECT_LE(fake_peer.maxAcceptedCount(), 2) + << "per-peer lane count exceeded in round " << round; +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + EXPECT_LE(maximum_observed_queue_depth.load(std::memory_order_acquire), + 64u); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 2u); + EXPECT_FALSE( + connecting_lane_had_current.load(std::memory_order_acquire)); +#endif + + fake_peer.closePeer(); + + std::vector final_states( + kRequestCount, TransferStatusEnum::WAITING); + bool status_query_failed = false; + bool all_terminal = false; + const auto deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(20); + while (!all_terminal && std::chrono::steady_clock::now() < deadline) { + all_terminal = true; + for (int i = 0; i < kRequestCount; ++i) { + if (final_states[i] == TransferStatusEnum::COMPLETED || + final_states[i] == TransferStatusEnum::FAILED) + continue; + + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + Status query = h.engine->getTransferStatus(batch_id, i, status); + if (!query.ok()) { + status_query_failed = true; + all_terminal = false; + continue; + } + final_states[i] = status.s; + if (status.s != TransferStatusEnum::COMPLETED && + status.s != TransferStatusEnum::FAILED) + all_terminal = false; + } + if (!all_terminal) + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + + EXPECT_FALSE(status_query_failed) << "round " << round; + EXPECT_TRUE(all_terminal) << "round " << round; + for (int i = 0; i < kRequestCount; ++i) { + EXPECT_NE(final_states[i], TransferStatusEnum::WAITING) + << "request " << i << " remained WAITING in round " << round; + EXPECT_TRUE(final_states[i] == TransferStatusEnum::COMPLETED || + final_states[i] == TransferStatusEnum::FAILED) + << "request " << i << " was not terminal in round " << round; + } + if (all_terminal) (void)h.engine->freeBatchID(batch_id); + } +} + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +TEST(TcpWriteVisibilityTest, + FailedLaneReconnectsWhileSiblingLaneRemainsUsable) { + constexpr int kRequestCount = 3; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "2"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedEnvVar status_timeout("MC_TCP_STATUS_TIMEOUT_SEC", "30"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17931", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + memset(h.pool, 0x6B, kLength); + const TransferRequest request = makeWriteRequest(h, kLength); + std::vector requests(kRequestCount, request); + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + + ASSERT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 2; }, + std::chrono::seconds(5))); + ASSERT_EQ(fake_peer.activeAcceptedCount(), 2); + + // Drop one busy lane while the sibling remains busy and therefore usable. + // The third request remains queued and must be pulled by a replacement + // connection for the failed lane, without entering peer-wide cooldown. + ASSERT_TRUE(fake_peer.closeOneAccepted()); + ASSERT_TRUE(waitForPredicate( + [] { return lane_terminal_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForAccepted(3, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connecting_count.load(std::memory_order_acquire) >= 3 && + lane_busy_count.load(std::memory_order_acquire) >= 3; + }, + std::chrono::seconds(5))); + + EXPECT_EQ(fake_peer.activeAcceptedCount(), 2); + EXPECT_EQ(cooldown_started_count.load(std::memory_order_acquire), 0); + EXPECT_GE(retry_armed_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(session_failure_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 2u); + + fake_peer.closePeer(); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), batch_id, kRequestCount, + std::chrono::seconds(10))); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + + hooks.reset(); + (void)h.engine->freeBatchID(batch_id); +} + +TEST(TcpWriteVisibilityTest, + ConnectFailedLaneRetriesWhileSiblingLaneRemainsUsable) { + constexpr int kRequestCount = 2; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "2"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedEnvVar status_timeout("MC_TCP_STATUS_TIMEOUT_SEC", "30"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/false, + /*block_retry=*/false, + /*block_retry_armed=*/false, + /*block_admission=*/false, + /*fail_second_connect=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17935", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + memset(h.pool, 0x6C, kLength); + const TransferRequest request = makeWriteRequest(h, kLength); + std::vector requests(kRequestCount, request); + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + + ASSERT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_injection_call_count.load( + std::memory_order_acquire) >= 3 && + lane_busy_count.load(std::memory_order_acquire) >= 2; + }, + std::chrono::seconds(5))); + + EXPECT_EQ(fake_peer.activeAcceptedCount(), 2); + EXPECT_EQ(lane_connect_injection_call_count.load(std::memory_order_acquire), + 3); + EXPECT_EQ(cooldown_started_count.load(std::memory_order_acquire), 0); + EXPECT_GE(retry_armed_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(connect_failure_count.load(std::memory_order_acquire), 0); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 2u); + + fake_peer.closePeer(); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), batch_id, kRequestCount, + std::chrono::seconds(10))); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + + hooks.reset(); + (void)h.engine->freeBatchID(batch_id); +} + +TEST(TcpWriteVisibilityTest, DirtyLastUsableLaneStartsFreshConnectRound) { + constexpr int kRequestCount = 2; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedEnvVar status_timeout("MC_TCP_STATUS_TIMEOUT_SEC", "30"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17936", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + memset(h.pool, 0x6E, kLength); + const TransferRequest request = makeWriteRequest(h, kLength); + const std::vector requests(kRequestCount, request); + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + + ASSERT_TRUE(fake_peer.waitForAccepted(1, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + ASSERT_EQ(fake_peer.activeAcceptedCount(), 1); + + // The only established lane fails while another request is still queued. + // Because this round had a successful connection, the queued request must + // get a fresh connect round instead of being failed as CONNECT_FAILED. + ASSERT_TRUE(fake_peer.closeOneAccepted()); + ASSERT_TRUE(waitForPredicate( + [] { return lane_terminal_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 2; }, + std::chrono::seconds(5))); + + EXPECT_EQ(fake_peer.activeAcceptedCount(), 1); + EXPECT_EQ(connect_failure_count.load(std::memory_order_acquire), 0); + EXPECT_GE(session_failure_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 1u); + + fake_peer.closePeer(); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), batch_id, kRequestCount, + std::chrono::seconds(10))); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + + hooks.reset(); + (void)h.engine->freeBatchID(batch_id); +} + +TEST(TcpWriteVisibilityTest, + ReconnectProbePrefersNeverTriedLanesOverFailedLane) { + constexpr int kRequestCount = 5; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "4"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedEnvVar status_timeout("MC_TCP_STATUS_TIMEOUT_SEC", "30"); + ScopedLaneHooks hooks; + tcpTransportSetLaneConnectFailureInjectionHookForTest( + failLaneOneAndTwoConnectAttempt); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17937", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + memset(h.pool, 0x6D, kLength); + const TransferRequest request = makeWriteRequest(h, kLength); + const std::vector requests(kRequestCount, request); + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + + // Lanes 1 and 2 fail on every probe. A starvation-prone selector retries + // lane 1 forever and never reaches the other retryable lane or lane 3. + // Each lane gets one probe per round, so lanes 0 and 3 are accepted. + ASSERT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_injection_call_count.load( + std::memory_order_acquire) >= 4; + }, + std::chrono::seconds(5))); + EXPECT_EQ(fake_peer.activeAcceptedCount(), 2); + EXPECT_GE(lane_connect_injection_call_count.load(std::memory_order_acquire), + 4); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 4u); + + fake_peer.closePeer(); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), batch_id, kRequestCount, + std::chrono::seconds(10))); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + + hooks.reset(); + (void)h.engine->freeBatchID(batch_id); +} +#endif + +TEST(TcpWriteVisibilityTest, + QueuedWorkAndBusyLanesCompleteExactlyOnceDuringShutdown) { + constexpr int kRequestCount = 16; + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "2"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "64"); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + ScopedLaneHooks hooks; +#endif + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17920", kLength); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = fake_peer.port(); + desc->tcp_proto_version = 2; + + memset(h.pool, 0x5A, kLength); + std::vector requests(kRequestCount); + for (auto& request : requests) { + request.opcode = TransferRequest::WRITE; + request.length = kLength; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + } + + auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + ASSERT_TRUE(fake_peer.waitForAccepted(2, std::chrono::seconds(5))); + EXPECT_LE(fake_peer.maxAcceptedCount(), 2); + + const auto start = std::chrono::steady_clock::now(); + h.engine.reset(); + const auto elapsed = std::chrono::steady_clock::now() - start; + EXPECT_LT(elapsed, std::chrono::seconds(5)); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_queue_depth.load(std::memory_order_acquire), + 64u); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 2u); + hooks.reset(); +#endif + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +TEST(TcpWriteVisibilityTest, + ConnectingLaneShutdownDetachesOwnershipAndIgnoresLateHandler) { + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17921", kLength); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = fake_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = kLength; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + auto batch_id = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, {request}).ok()); + + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_handler_entered.load(std::memory_order_acquire); + }, + std::chrono::seconds(5))); + EXPECT_FALSE(connecting_lane_had_current.load(std::memory_order_acquire)); + + auto engine = std::move(h.engine); + auto destruction = + std::async(std::launch::async, + [engine = std::move(engine)]() mutable { engine.reset(); }); + + // Let shutdown invalidate the lane epoch before the blocked resolve + // handler returns. io_context::stop() cannot interrupt that handler, so + // destruction must wait for this explicit release and then join it. + std::this_thread::sleep_for(std::chrono::milliseconds(100)); + EXPECT_EQ(destruction.wait_for(std::chrono::milliseconds(0)), + std::future_status::timeout); + releaseLaneConnectHandler(); + + ASSERT_EQ(destruction.wait_for(std::chrono::seconds(5)), + std::future_status::ready); + destruction.get(); + + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + EXPECT_GE(late_lane_handler_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 1u); + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} + +TEST(TcpWriteVisibilityTest, + ReconnectRoundsAreRateLimitedAndCooldownQueueIsBounded) { + constexpr int kCooldownRequestCount = 4; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "3"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "500"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + UnavailableTcpPeer unavailable_peer; + ASSERT_TRUE(unavailable_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17925", 64 * 1024); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = unavailable_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + + const auto first_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(first_batch, {request}).ok()); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), first_batch, 1, + std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return cooldown_started_count.load(std::memory_order_acquire) >= 1; + }, + std::chrono::seconds(2))); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 1); + + std::vector cooldown_requests(kCooldownRequestCount, + request); + const auto cooldown_batch = + h.engine->allocateBatchID(kCooldownRequestCount); + ASSERT_TRUE( + h.engine->submitTransfer(cooldown_batch, cooldown_requests).ok()); + ASSERT_TRUE(waitForPredicate( + [] { return retry_armed_count.load(std::memory_order_acquire) == 1; }, + std::chrono::seconds(2))); + + // The first three requests enter the work queue; the fourth waits in the + // pending-admission FIFO instead of failing synchronously. No second + // connection round starts before the retry timer. + for (int task_id = 0; task_id < 3; ++task_id) { + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + ASSERT_TRUE( + h.engine->getTransferStatus(cooldown_batch, task_id, status).ok()); + EXPECT_EQ(status.s, TransferStatusEnum::WAITING); + } + TransferStatus overflow_status; + overflow_status.s = TransferStatusEnum::WAITING; + ASSERT_TRUE( + h.engine->getTransferStatus(cooldown_batch, 3, overflow_status).ok()); + EXPECT_EQ(overflow_status.s, TransferStatusEnum::WAITING); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_pending_count.load(std::memory_order_acquire), 1); + EXPECT_LE(maximum_observed_queue_depth.load(std::memory_order_acquire), 3u); + ASSERT_TRUE(waitForPredicate( + [] { + return queue_timeout_failure_count.load( + std::memory_order_acquire) == 1; + }, + std::chrono::seconds(2))); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(retry_armed_count.load(std::memory_order_acquire), 1); + + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), cooldown_batch, + kCooldownRequestCount, + std::chrono::seconds(5))); + EXPECT_EQ(retry_fired_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(retry_armed_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 2); + EXPECT_EQ(connect_failure_count.load(std::memory_order_acquire), 4); + expectEverySliceCompletedExactlyOnceAfterShutdown(first_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(cooldown_batch); + + hooks.reset(); + (void)h.engine->freeBatchID(first_batch); + (void)h.engine->freeBatchID(cooldown_batch); +} + +TEST(TcpWriteVisibilityTest, + RetryTimerSerializesDelayedPumpAndCooldownTransition) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "2"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/false, + /*block_retry=*/false, + /*block_retry_armed=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + UnavailableTcpPeer unavailable_peer; + ASSERT_TRUE(unavailable_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17928", 64 * 1024); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = unavailable_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + + const auto exhausted_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(exhausted_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return cooldown_started_count.load(std::memory_order_acquire) >= + 1 && + connect_failure_count.load(std::memory_order_acquire) >= 1; + }, + std::chrono::seconds(2))); + + const auto first_pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(first_pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return retry_armed_observer_entered.load(std::memory_order_acquire); + }, + std::chrono::seconds(2))); + + // The worker is blocked after timer registration and outside the group + // mutex. This admission posts a pump before the deadline; keeping the + // worker blocked past expiry makes that pump run before the ready timer. + const auto second_pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(second_pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return maximum_observed_queue_depth.load( + std::memory_order_acquire) == 2; + }, + std::chrono::seconds(2))); + std::this_thread::sleep_for(std::chrono::milliseconds(1200)); + + EXPECT_EQ(retry_armed_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(queue_rejection_count.load(std::memory_order_acquire), 0); + EXPECT_LE(maximum_observed_queue_depth.load(std::memory_order_acquire), 2u); + + releaseRetryArmedObserver(); + ASSERT_TRUE(waitForPredicate( + [] { + return retry_fired_count.load(std::memory_order_acquire) == 1 && + connect_failure_count.load(std::memory_order_acquire) == 3; + }, + std::chrono::seconds(5))); + + h.engine.reset(); + expectEverySliceCompletedExactlyOnceAfterShutdown(exhausted_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(first_pending_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(second_pending_batch); + EXPECT_EQ(retry_armed_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(retry_fired_count.load(std::memory_order_acquire), 1); + // One initial attempt plus one timer-owned retry proves the delayed pump + // neither reset an active round nor created an extra immediate attempt. + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 2); + EXPECT_EQ(connect_failure_count.load(std::memory_order_acquire), 3); + EXPECT_LE(maximum_observed_queue_depth.load(std::memory_order_acquire), 2u); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(exhausted_batch); + reclaimBatchDescAfterEngineShutdownForTest(first_pending_batch); + reclaimBatchDescAfterEngineShutdownForTest(second_pending_batch); +} + +TEST(TcpWriteVisibilityTest, + ShutdownWithPendingRetryTimerCompletesAcceptedWorkOnce) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "4"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + UnavailableTcpPeer unavailable_peer; + ASSERT_TRUE(unavailable_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17927", 64 * 1024); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = unavailable_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + + const auto exhausted_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(exhausted_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return cooldown_started_count.load(std::memory_order_acquire) >= + 1 && + connect_failure_count.load(std::memory_order_acquire) >= 1; + }, + std::chrono::seconds(2))); + + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { return retry_armed_count.load(std::memory_order_acquire) == 1; }, + std::chrono::seconds(2))); + + const auto shutdown_start = std::chrono::steady_clock::now(); + h.engine.reset(); + EXPECT_LT(std::chrono::steady_clock::now() - shutdown_start, + std::chrono::seconds(5)); + + expectEverySliceCompletedExactlyOnceAfterShutdown(exhausted_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(shutdown_failure_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(retry_fired_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(exhausted_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, + ShutdownInvalidatesPendingRetryAndLateHandlerCannotReconnect) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "4"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/false, + /*block_retry=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + UnavailableTcpPeer unavailable_peer; + ASSERT_TRUE(unavailable_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17926", 64 * 1024); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = unavailable_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + + const auto exhausted_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(exhausted_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return cooldown_started_count.load(std::memory_order_acquire) >= + 1 && + connect_failure_count.load(std::memory_order_acquire) >= 1; + }, + std::chrono::seconds(2))); + + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { return retry_armed_count.load(std::memory_order_acquire) == 1; }, + std::chrono::seconds(2))); + ASSERT_TRUE(waitForPredicate( + [] { return retry_handler_entered.load(std::memory_order_acquire); }, + std::chrono::seconds(3))); + + auto engine = std::move(h.engine); + auto destruction = + std::async(std::launch::async, + [engine = std::move(engine)]() mutable { engine.reset(); }); + ASSERT_TRUE(waitForPredicate( + [] { + return shutdown_failure_count.load(std::memory_order_acquire) == 1; + }, + std::chrono::seconds(2))); + EXPECT_EQ(destruction.wait_for(std::chrono::milliseconds(0)), + std::future_status::timeout); + + const int connects_before_release = + lane_connecting_count.load(std::memory_order_acquire); + releaseRetryHandler(); + ASSERT_EQ(destruction.wait_for(std::chrono::seconds(5)), + std::future_status::ready); + destruction.get(); + + expectEverySliceCompletedExactlyOnceAfterShutdown(exhausted_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(lane_connecting_count.load(std::memory_order_acquire), + connects_before_release); + EXPECT_EQ(retry_fired_count.load(std::memory_order_acquire), 0); + EXPECT_GE(retry_late_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(exhausted_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, + TaskGroupShutdownCompletesNotYetStartedSlicesExactlyOnce) { + constexpr int kRequestCount = 3; + constexpr size_t kLength = 64 * 1024; + + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedLaneHooks hooks; + + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17936", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + memset(h.pool, 0x6D, kLength); + TransferRequest request = makeWriteRequest(h, kLength); + request.task_group_id = 1; + + std::vector requests(kRequestCount, request); + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + + // Task-group sequencing starts only the first Slice. Hold it BUSY so the + // remaining Slices have been created but have not entered the lane queue. + ASSERT_TRUE(fake_peer.waitForAccepted(1, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + + h.engine.reset(); + + // Shutdown must not strand the sequence tail merely because the TCP + // io_context has already been stopped. Every pre-created Slice must become + // terminal exactly once. + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + + const auto& batch = Transport::toBatchDesc(batch_id); + ASSERT_EQ(batch.task_list.size(), kRequestCount); + for (size_t task_id = 0; task_id < batch.task_list.size(); ++task_id) { + const auto& task = batch.task_list[task_id]; + const uint64_t success = + __atomic_load_n(&task.success_slice_count, __ATOMIC_RELAXED); + const uint64_t failed = + __atomic_load_n(&task.failed_slice_count, __ATOMIC_RELAXED); + const uint64_t slices = + __atomic_load_n(&task.slice_count, __ATOMIC_RELAXED); + + EXPECT_EQ(slices, 1u) << "task " << task_id; + EXPECT_EQ(success, 0u) << "task " << task_id; + EXPECT_EQ(failed, 1u) << "task " << task_id; + ASSERT_EQ(task.slice_list.size(), 1u); + EXPECT_EQ(task.slice_list.front()->status, Transport::Slice::FAILED) + << "task " << task_id; + } + + EXPECT_EQ(shutdown_failure_count.load(std::memory_order_acquire), + kRequestCount); + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} + +TEST(TcpWriteVisibilityTest, LaneTerminalTypesAreMoveOnly) { + EXPECT_TRUE(tcpTransportLaneTypesAreMoveOnlyForTest()); +} +#endif + +TEST(TcpWriteVisibilityTest, OneLaneReusesCleanSocketInFifoOrder) { + constexpr int kRequestCount = 8; + constexpr size_t kLength = 64 * 1024; + constexpr size_t kPoolSize = kRequestCount * kLength; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "16"); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + ReusingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17922", kPoolSize); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = fake_peer.port(); + desc->tcp_proto_version = 2; + + memset(h.pool, 0x6B, kLength); + std::vector requests(kRequestCount); + std::vector expected_addresses; + expected_addresses.reserve(kRequestCount); + for (int i = 0; i < kRequestCount; ++i) { + auto& request = requests[i]; + request.opcode = TransferRequest::WRITE; + request.length = kLength; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base + i * kLength; + expected_addresses.push_back(request.target_offset); + } + + auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + ASSERT_TRUE( + fake_peer.waitForRequests(kRequestCount, std::chrono::seconds(10))); + + for (int i = 0; i < kRequestCount; ++i) { + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + const auto deadline = + std::chrono::steady_clock::now() + std::chrono::seconds(5); + while (status.s == TransferStatusEnum::WAITING && + std::chrono::steady_clock::now() < deadline) { + ASSERT_TRUE(h.engine->getTransferStatus(batch_id, i, status).ok()); + if (status.s == TransferStatusEnum::WAITING) + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + EXPECT_EQ(status.s, TransferStatusEnum::COMPLETED) << "request " << i; + } + + EXPECT_EQ(fake_peer.acceptedCount(), 1); + EXPECT_EQ(fake_peer.requestAddresses(), expected_addresses); + (void)h.engine->freeBatchID(batch_id); +} + +#ifdef MOONCAKE_TCP_TRANSPORT_TEST_HOOKS +TEST(TcpWriteVisibilityTest, + PendingAdmissionWaitsInsteadOfFailingSynchronously) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "2"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17929", 64 * 1024); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + const auto request = makeWriteRequest(h, 1); + const auto batch_id = h.engine->allocateBatchID(2); + const auto submit_start = std::chrono::steady_clock::now(); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, {request, request}).ok()); + const auto submit_elapsed = std::chrono::steady_clock::now() - submit_start; + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_handler_entered.load( + std::memory_order_acquire) && + admission_pending_count.load(std::memory_order_acquire) == 1; + }, + std::chrono::seconds(5))); + + TransferStatus pending_status; + pending_status.s = TransferStatusEnum::WAITING; + ASSERT_TRUE(h.engine->getTransferStatus(batch_id, 1, pending_status).ok()); + EXPECT_EQ(pending_status.s, TransferStatusEnum::WAITING); + EXPECT_LT(submit_elapsed, std::chrono::seconds(2)); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(maximum_observed_queue_depth.load(std::memory_order_acquire), 1u); + EXPECT_EQ(maximum_observed_pending_depth.load(std::memory_order_acquire), + 1u); + + releaseLaneConnectHandler(); + h.engine.reset(); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} + +TEST(TcpWriteVisibilityTest, PendingAdmissionTimesOutExactlyOnce) { + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "100"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17930", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + const auto request = makeWriteRequest(h, kLength); + const auto active_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(active_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForAccepted(1, std::chrono::seconds(5))); + + const auto queued_batch = h.engine->allocateBatchID(1); + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(queued_batch, {request}).ok()); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return queue_timeout_failure_count.load( + std::memory_order_acquire) == 1; + }, + std::chrono::seconds(5))); + + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_timer_fired_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(admission_promoted_count.load(std::memory_order_acquire), 0); + h.engine.reset(); + expectEverySliceCompletedExactlyOnceAfterShutdown(active_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(queued_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(active_batch); + reclaimBatchDescAfterEngineShutdownForTest(queued_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, PendingAdmissionsPromoteInFifoOrder) { + constexpr int kRequestCount = 4; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "3"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + ReusingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17931", 64 * 1024); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + std::vector requests; + std::vector expected_addresses; + for (int i = 0; i < kRequestCount; ++i) { + const uint64_t address = h.remote_base + static_cast(i); + requests.push_back(makeWriteRequest(h, 1, address)); + expected_addresses.push_back(address); + } + + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_handler_entered.load( + std::memory_order_acquire) && + admission_pending_count.load(std::memory_order_acquire) == 3; + }, + std::chrono::seconds(5))); + EXPECT_EQ(maximum_observed_queue_depth.load(std::memory_order_acquire), 1u); + EXPECT_EQ(maximum_observed_pending_depth.load(std::memory_order_acquire), + 3u); + + releaseLaneConnectHandler(); + ASSERT_TRUE( + fake_peer.waitForRequests(kRequestCount, std::chrono::seconds(10))); + ASSERT_TRUE(waitForBatchTerminal(h.engine.get(), batch_id, kRequestCount, + std::chrono::seconds(5))); + EXPECT_EQ(fake_peer.requestAddresses(), expected_addresses); + EXPECT_GE(admission_promoted_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + + hooks.reset(); + (void)h.engine->freeBatchID(batch_id); +} + +TEST(TcpWriteVisibilityTest, PendingAdmissionHardBoundRejectsImmediately) { + constexpr int kRequestCount = 3; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17923", 64 * 1024); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + std::vector requests(kRequestCount); + for (auto& request : requests) { + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + } + + const auto batch_id = h.engine->allocateBatchID(kRequestCount); + ASSERT_TRUE(h.engine->submitTransfer(batch_id, requests).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_handler_entered.load(std::memory_order_acquire); + }, + std::chrono::seconds(5))); + + EXPECT_EQ(queue_rejection_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(admission_hard_rejection_count.load(std::memory_order_acquire), + 1); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(connect_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(maximum_observed_queue_depth.load(std::memory_order_acquire), 1u); + EXPECT_EQ(maximum_observed_pending_depth.load(std::memory_order_acquire), + 1u); + EXPECT_FALSE(connecting_lane_had_current.load(std::memory_order_acquire)); + for (int task_id = 0; task_id < kRequestCount; ++task_id) { + TransferStatus status; + status.s = TransferStatusEnum::WAITING; + ASSERT_TRUE( + h.engine->getTransferStatus(batch_id, task_id, status).ok()); + EXPECT_EQ(status.s, task_id == kRequestCount - 1 + ? TransferStatusEnum::FAILED + : TransferStatusEnum::WAITING); + } + + releaseLaneConnectHandler(); + h.engine.reset(); + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} + +TEST(TcpWriteVisibilityTest, + ShutdownWithPendingAdmissionsAndArmedTimerCompletesOnce) { + constexpr size_t kLength = 64 * 1024; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks; + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17932", kLength); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + const auto request = makeWriteRequest(h, kLength); + const auto active_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(active_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { return lane_busy_count.load(std::memory_order_acquire) >= 1; }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForAccepted(1, std::chrono::seconds(5))); + + const auto queued_batch = h.engine->allocateBatchID(1); + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(queued_batch, {request}).ok()); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_pending_count.load(std::memory_order_acquire) == + 1 && + admission_timer_armed_count.load( + std::memory_order_acquire) == 1; + }, + std::chrono::seconds(5))); + + const auto shutdown_start = std::chrono::steady_clock::now(); + h.engine.reset(); + EXPECT_LT(std::chrono::steady_clock::now() - shutdown_start, + std::chrono::seconds(5)); + expectEverySliceCompletedExactlyOnceAfterShutdown(active_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(queued_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(shutdown_failure_count.load(std::memory_order_acquire), 3); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_promoted_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_timer_fired_count.load(std::memory_order_acquire), 0); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(active_batch); + reclaimBatchDescAfterEngineShutdownForTest(queued_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, + AdmissionTimeoutWinsBeforeCapacityReleasePreservesOneOwner) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "100"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/false, + /*block_retry=*/false, + /*block_retry_armed=*/false, + /*block_admission=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + ReusingWriteServer fake_peer(/*hold_first_ack=*/true); + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17933", 64 * 1024); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + const auto request = makeWriteRequest(h, 1); + const auto active_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(active_batch, {request}).ok()); + ASSERT_TRUE(fake_peer.waitForFirstRequest(std::chrono::seconds(5))); + + const auto queued_batch = h.engine->allocateBatchID(1); + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(queued_batch, {request}).ok()); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_handler_entered.load(std::memory_order_acquire); + }, + std::chrono::seconds(5))); + + // The expired timer handler is paused outside the group mutex. Make the + // active session's capacity release ready, then let the timer and queued + // completion contend in a deterministic order on the single io worker. + fake_peer.releaseFirstAck(); + ASSERT_TRUE(fake_peer.waitForFirstAckSent(std::chrono::seconds(5))); + releaseAdmissionHandler(); + ASSERT_TRUE(waitForPredicate( + [] { + return queue_timeout_failure_count.load( + std::memory_order_acquire) == 1; + }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForRequests(2, std::chrono::seconds(5))); + + h.engine.reset(); + expectEverySliceCompletedExactlyOnceAfterShutdown(active_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(queued_batch); + expectEverySliceCompletedExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_promoted_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_timer_fired_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(active_batch); + reclaimBatchDescAfterEngineShutdownForTest(queued_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, + CapacityReleasePromotesBeforeDeadlineAndCancelledTimerIsLate) { + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "1"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "1"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/false, + /*block_retry=*/false, + /*block_retry_armed=*/false, + /*block_admission=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + ReusingWriteServer fake_peer(/*hold_first_ack=*/true); + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17934", 64 * 1024); + ASSERT_TRUE(h.ok); + pointTcpSegmentAt(h, fake_peer.port()); + + const auto request = makeWriteRequest(h, 1); + const auto active_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(active_batch, {request}).ok()); + ASSERT_TRUE(fake_peer.waitForFirstRequest(std::chrono::seconds(5))); + + const auto queued_batch = h.engine->allocateBatchID(1); + const auto pending_batch = h.engine->allocateBatchID(1); + ASSERT_TRUE(h.engine->submitTransfer(queued_batch, {request}).ok()); + ASSERT_TRUE(h.engine->submitTransfer(pending_batch, {request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_pending_count.load(std::memory_order_acquire) == + 1 && + admission_timer_armed_count.load( + std::memory_order_acquire) == 1; + }, + std::chrono::seconds(5))); + + // Completing the active exchange frees a work-queue slot before the long + // deadline. The pump promotes the pending item, invalidates the timer + // epoch, and cancels the now-unneeded shared admission timer. + fake_peer.releaseFirstAck(); + ASSERT_TRUE(fake_peer.waitForFirstAckSent(std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_promoted_count.load(std::memory_order_acquire) >= + 1; + }, + std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_handler_entered.load(std::memory_order_acquire); + }, + std::chrono::seconds(5))); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(admission_timer_fired_count.load(std::memory_order_acquire), 0); + + releaseAdmissionHandler(); + ASSERT_TRUE(waitForPredicate( + [] { + return admission_timer_late_count.load(std::memory_order_acquire) >= + 1; + }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForRequests(3, std::chrono::seconds(5))); + ASSERT_TRUE(waitForPredicate( + [] { return lane_terminal_count.load(std::memory_order_acquire) >= 3; }, + std::chrono::seconds(5))); + + h.engine.reset(); + expectEverySliceSucceededExactlyOnceAfterShutdown(active_batch); + expectEverySliceSucceededExactlyOnceAfterShutdown(queued_batch); + expectEverySliceSucceededExactlyOnceAfterShutdown(pending_batch); + EXPECT_EQ(queue_timeout_failure_count.load(std::memory_order_acquire), 0); + EXPECT_EQ(queue_full_failure_count.load(std::memory_order_acquire), 0); + EXPECT_GE(admission_promoted_count.load(std::memory_order_acquire), 1); + EXPECT_EQ(admission_timer_fired_count.load(std::memory_order_acquire), 0); + EXPECT_GE(admission_timer_late_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(active_batch); + reclaimBatchDescAfterEngineShutdownForTest(queued_batch); + reclaimBatchDescAfterEngineShutdownForTest(pending_batch); +} + +TEST(TcpWriteVisibilityTest, + ConcurrentAdmissionShutdownStressPreservesLaneOwnership) { + constexpr size_t kSubmissionCount = 40; + constexpr size_t kSubmitterCount = 4; + constexpr size_t kQueueCapacity = 8; + ScopedEnvVar lanes("MC_TCP_LANES_PER_PEER", "2"); + ScopedEnvVar queue_capacity("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", "8"); + ScopedEnvVar pending_capacity("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", + "1"); + ScopedEnvVar admission_timeout("MC_TCP_ADMISSION_TIMEOUT_MS", "10000"); + ScopedLaneHooks hooks(/*block_first_connect_handler=*/false, + /*block_after_busy=*/true); + const char* env = std::getenv("MC_METADATA_SERVER"); + const std::string metadata_server = env ? env : "P2PHANDSHAKE"; + + HoldingWriteServer fake_peer; + ASSERT_TRUE(fake_peer.ok()); + + EngineHandle h; + h.init(metadata_server, "127.0.0.2:17924", 64 * 1024); + ASSERT_TRUE(h.ok); + + auto desc = h.engine->getMetadata()->getSegmentDescByID(h.segment_id); + ASSERT_NE(desc, nullptr); + desc->tcp_data_port = fake_peer.port(); + desc->tcp_proto_version = 2; + + TransferRequest request; + request.opcode = TransferRequest::WRITE; + request.length = 1; + request.source = h.pool; + request.target_id = h.segment_id; + request.target_offset = h.remote_base; + + const auto busy_batch_id = h.engine->allocateBatchID(2); + ASSERT_TRUE( + h.engine->submitTransfer(busy_batch_id, {request, request}).ok()); + ASSERT_TRUE(waitForPredicate( + [] { + return lane_connect_handler_entered.load( + std::memory_order_acquire) && + lane_busy_count.load(std::memory_order_acquire) >= 1; + }, + std::chrono::seconds(5))); + ASSERT_TRUE(fake_peer.waitForAccepted(1, std::chrono::seconds(5))); + + std::vector batch_ids(kSubmissionCount); + for (auto& batch_id : batch_ids) batch_id = h.engine->allocateBatchID(1); + std::vector> submit_results(kSubmissionCount); + for (auto& result : submit_results) result.store(0); + + std::vector submitters; + for (size_t thread_id = 0; thread_id < kSubmitterCount; ++thread_id) { + submitters.emplace_back([&, thread_id] { + for (size_t i = thread_id; i < kSubmissionCount; + i += kSubmitterCount) { + submit_results[i].store( + h.engine->submitTransfer(batch_ids[i], {request}).ok() ? 1 + : -1, + std::memory_order_release); + } + }); + } + for (auto& submitter : submitters) submitter.join(); + for (const auto& result : submit_results) + EXPECT_EQ(result.load(std::memory_order_acquire), 1); + + EXPECT_EQ(maximum_observed_queue_depth.load(std::memory_order_acquire), + kQueueCapacity); + EXPECT_LE(maximum_observed_socket_count.load(std::memory_order_acquire), + 2u); + EXPECT_FALSE(connecting_lane_had_current.load(std::memory_order_acquire)); + EXPECT_EQ(queue_rejection_count.load(std::memory_order_acquire), + static_cast(kSubmissionCount - (kQueueCapacity - 1) - 1)); + EXPECT_LE(maximum_observed_pending_depth.load(std::memory_order_acquire), + 1u); + + auto engine = std::move(h.engine); + auto destruction = + std::async(std::launch::async, + [engine = std::move(engine)]() mutable { engine.reset(); }); + std::this_thread::sleep_for(std::chrono::milliseconds(100)); + releaseLaneConnectHandler(); + ASSERT_EQ(destruction.wait_for(std::chrono::seconds(5)), + std::future_status::ready); + destruction.get(); + + expectEverySliceCompletedExactlyOnceAfterShutdown(busy_batch_id); + for (const auto batch_id : batch_ids) + expectEverySliceCompletedExactlyOnceAfterShutdown(batch_id); + EXPECT_EQ(lane_shutdown_clean_count.load(std::memory_order_acquire), 1); + EXPECT_GE(late_lane_handler_count.load(std::memory_order_acquire), 1); + + hooks.reset(); + reclaimBatchDescAfterEngineShutdownForTest(busy_batch_id); + for (const auto batch_id : batch_ids) + reclaimBatchDescAfterEngineShutdownForTest(batch_id); +} +#endif From 54e85517f6f30bd1916e159ee892906d7448506d Mon Sep 17 00:00:00 2001 From: Aoi Date: Tue, 11 Aug 2026 11:06:50 +0800 Subject: [PATCH 022/483] [Doc] Align Python setup API documentation (#3357) --- .../api-reference/python/mooncake-store.md | 38 ++++++++++++++----- 1 file changed, 29 insertions(+), 9 deletions(-) diff --git a/docs/source/api-reference/python/mooncake-store.md b/docs/source/api-reference/python/mooncake-store.md index dd90000cbe..96094b015a 100644 --- a/docs/source/api-reference/python/mooncake-store.md +++ b/docs/source/api-reference/python/mooncake-store.md @@ -1054,30 +1054,50 @@ def setup( self, local_hostname: str, metadata_server: str, - global_segment_size: int = 16777216, - local_buffer_size: int = 1073741824, - protocol: str = "tcp", - rdma_devices: str = "", + global_segment_size: int, + local_buffer_size: int, + protocol: str, + rdma_devices: str, master_server_addr: str, engine: Optional[TransferEngine] = None, enable_ssd_offload: bool = False, ssd_offload_path: str = "", tenant_id: str = "default", + enable_client_http_server: bool = False, + client_http_port: int = 9300, ) -> int ``` +The positional overload requires every argument through +`master_server_addr`. To use defaults for those fields, pass a configuration +dictionary instead: + +```python +def setup(self, config: Dict[str, object]) -> int +``` + +The dictionary overload requires `local_hostname` and `metadata_server`. Its +other keys are optional; the defaults are `16777216` (16 MiB) for both +`global_segment_size` and `local_buffer_size`, `"tcp"` for `protocol`, an empty +string for `rdma_devices`, and `"127.0.0.1:50051"` for +`master_server_addr`. It also accepts `ipc_socket_path` and the optional +configuration fields listed below. The `engine` argument is available only in +the positional overload. + **Parameters:** - `local_hostname` (str): **Required**. Local hostname and port (e.g., "localhost" or "localhost:12345") - `metadata_server` (str): **Required**. Metadata connection string, e.g. `"P2PHANDSHAKE"` or `"http://localhost:8080/metadata"`. -- `global_segment_size` (int): Memory segment size in bytes for mounting. -- `local_buffer_size` (int): Local buffer size in bytes. -- `protocol` (str): Network protocol, usually `"tcp"`, `"rdma"`, `"efa"`, `"cxl"`, or `"ascend"` depending on the build. -- `rdma_devices` (str): RDMA/EFA device name(s), e.g. `"mlx5_0"` or `"mlx5_0,mlx5_1"`. Leave empty to auto-discover NICs unless `MC_MS_AUTO_DISC=0`; always empty for TCP. -- `master_server_addr` (str): **Required**. Master server address (e.g., "localhost:50051") +- `global_segment_size` (int): **Required by the positional overload**. Memory segment size in bytes for mounting. +- `local_buffer_size` (int): **Required by the positional overload**. Local buffer size in bytes. +- `protocol` (str): **Required by the positional overload**. Network protocol, usually `"tcp"`, `"rdma"`, `"efa"`, `"cxl"`, or `"ascend"` depending on the build. +- `rdma_devices` (str): **Required by the positional overload**. RDMA/EFA device name(s), e.g. `"mlx5_0"` or `"mlx5_0,mlx5_1"`. Leave empty to auto-discover NICs unless `MC_MS_AUTO_DISC=0`; always empty for TCP. +- `master_server_addr` (str): **Required by the positional overload**. Master server address (e.g., "localhost:50051") - `engine` (Optional[TransferEngine]): Existing Transfer Engine instance to reuse. Defaults to `None`. - `enable_ssd_offload` (bool): Enable client-side SSD offload support. Defaults to `False`. - `ssd_offload_path` (str): SSD offload directory. When provided, overrides the storage path environment configuration. - `tenant_id` (str): Tenant namespace for object keys. Defaults to `"default"`. +- `enable_client_http_server` (bool): Enable the client-local `/health`, `/metrics`, and `/metrics/summary` HTTP endpoints. Defaults to `False`. +- `client_http_port` (int): Port for the client-local HTTP endpoints. Defaults to `9300`. **Store segment pinned memory:** CUDA-enabled builds can register Store-managed host segments as pinned memory when `MC_STORE_PIN_MEMORY_MAX_BYTES` is set to a From ca7695dca469c72d8d7e306bc5e9cc572fe3c34a Mon Sep 17 00:00:00 2001 From: Aoi Date: Tue, 11 Aug 2026 11:07:08 +0800 Subject: [PATCH 023/483] [CI/Build] Isolate nightly Python tests in venv (#3373) --- .github/workflows/nightly.yml | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/.github/workflows/nightly.yml b/.github/workflows/nightly.yml index 6976f4c8d5..0acdb377c4 100644 --- a/.github/workflows/nightly.yml +++ b/.github/workflows/nightly.yml @@ -306,10 +306,13 @@ jobs: - name: Build and install Python wheel for integration tests run: | + python -m venv "${RUNNER_TEMP}/mooncake-nightly-venv" + echo "${RUNNER_TEMP}/mooncake-nightly-venv/bin" >> "$GITHUB_PATH" + source "${RUNNER_TEMP}/mooncake-nightly-venv/bin/activate" export LD_LIBRARY_PATH=${LD_LIBRARY_PATH:-}:/usr/local/lib export CUDA_HOME=/usr/local/cuda PYTHON_VERSION=3.12 OUTPUT_DIR=dist ./scripts/build_wheel.sh - pip install mooncake-wheel/dist/*.whl + python -m pip install mooncake-wheel/dist/*.whl - name: Run Python integration tests (full suite) env: From 38924b22860ba239b5f285bcb525455c02ce61f8 Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Tue, 11 Aug 2026 11:32:00 +0800 Subject: [PATCH 024/483] [Store] Add batch OpLog snapshot metadata protocol (#3178) Co-authored-by: Yuchen Kou --- .../ha/snapshot/batch_oplog/metadata.h | 86 ++++ mooncake-store/src/CMakeLists.txt | 1 + .../src/ha/snapshot/batch_oplog/metadata.cpp | 409 ++++++++++++++++++ mooncake-store/tests/CMakeLists.txt | 2 + .../ha/snapshot/batch_oplog/metadata_test.cpp | 267 ++++++++++++ 5 files changed, 765 insertions(+) create mode 100644 mooncake-store/include/ha/snapshot/batch_oplog/metadata.h create mode 100644 mooncake-store/src/ha/snapshot/batch_oplog/metadata.cpp create mode 100644 mooncake-store/tests/ha/snapshot/batch_oplog/metadata_test.cpp diff --git a/mooncake-store/include/ha/snapshot/batch_oplog/metadata.h b/mooncake-store/include/ha/snapshot/batch_oplog/metadata.h new file mode 100644 index 0000000000..1c2c701bad --- /dev/null +++ b/mooncake-store/include/ha/snapshot/batch_oplog/metadata.h @@ -0,0 +1,86 @@ +#pragma once + +#include +#include +#include +#include + +#include + +#include "types.h" + +namespace mooncake::ha { + +inline constexpr uint32_t kBatchOpLogSnapshotSchemaVersion = 1; +inline constexpr char kBatchOpLogSnapshotFormat[] = + "standby-oplog-materialized/v1"; + +struct BatchOpLogSnapshotDescriptor { + uint32_t schema_version{kBatchOpLogSnapshotSchemaVersion}; + std::string snapshot_format{kBatchOpLogSnapshotFormat}; + std::string snapshot_id; + uint64_t last_included_seq{0}; + uint64_t last_included_batch_id{0}; + ViewVersionId producer_view_version{0}; + std::string manifest_key; + uint64_t manifest_size{0}; + uint32_t manifest_crc32c{0}; + int64_t created_at_ms{0}; +}; + +struct BatchOpLogSnapshotObjectDescriptor { + std::string key; + uint64_t stored_size{0}; + uint32_t crc32c{0}; +}; + +struct BatchOpLogSnapshotChunkDescriptor { + uint64_t chunk_index{0}; + std::string key; + uint64_t object_count{0}; + uint64_t stored_size{0}; + uint32_t crc32c{0}; +}; + +struct BatchOpLogSnapshotManifest { + uint32_t schema_version{kBatchOpLogSnapshotSchemaVersion}; + std::string snapshot_format{kBatchOpLogSnapshotFormat}; + std::string snapshot_id; + uint64_t last_included_seq{0}; + uint64_t last_included_batch_id{0}; + ViewVersionId producer_view_version{0}; + BatchOpLogSnapshotObjectDescriptor segments; + std::vector object_chunks; +}; + +std::string EncodeBatchOpLogSnapshotDescriptor( + const BatchOpLogSnapshotDescriptor& descriptor); +tl::expected +DecodeBatchOpLogSnapshotDescriptor(std::string_view value); + +std::string EncodeBatchOpLogSnapshotManifest( + const BatchOpLogSnapshotManifest& manifest); +tl::expected +DecodeBatchOpLogSnapshotManifest(std::string_view value); + +std::string BuildBatchOpLogSnapshotId(uint64_t last_included_batch_id, + int64_t maintenance_lease_id); + +std::string BuildBatchOpLogSnapshotMaintenanceKey( + const std::string& cluster_id); +std::string BuildBatchOpLogSnapshotLatestKey(const std::string& cluster_id); +std::string BuildBatchOpLogSnapshotFallbackKey(const std::string& cluster_id); +std::string BuildBatchOpLogSnapshotCompactionFloorKey( + const std::string& cluster_id); + +std::string BuildBatchOpLogSnapshotDescriptorKey( + const std::string& snapshot_root, std::string_view snapshot_id); +std::string BuildBatchOpLogSnapshotManifestKey(const std::string& snapshot_root, + std::string_view snapshot_id); +std::string BuildBatchOpLogSnapshotSegmentsKey(const std::string& snapshot_root, + std::string_view snapshot_id); +std::string BuildBatchOpLogSnapshotObjectChunkKey( + const std::string& snapshot_root, std::string_view snapshot_id, + uint64_t chunk_index); + +} // namespace mooncake::ha diff --git a/mooncake-store/src/CMakeLists.txt b/mooncake-store/src/CMakeLists.txt index f20c7cbf7d..6fbf2f1a9a 100644 --- a/mooncake-store/src/CMakeLists.txt +++ b/mooncake-store/src/CMakeLists.txt @@ -51,6 +51,7 @@ set(MOONCAKE_STORE_SOURCES ha/kv/etcd_ha_kv_backend.cpp ha/standby_controller.cpp ha/snapshot/catalog_backed_snapshot_provider.cpp + ha/snapshot/batch_oplog/metadata.cpp ha/snapshot/master_snapshot_codec.cpp ha/snapshot/object/snapshot_object_store.cpp ha/snapshot/object/backends/local/local_file_snapshot_object_store.cpp diff --git a/mooncake-store/src/ha/snapshot/batch_oplog/metadata.cpp b/mooncake-store/src/ha/snapshot/batch_oplog/metadata.cpp new file mode 100644 index 0000000000..e99a03f8b0 --- /dev/null +++ b/mooncake-store/src/ha/snapshot/batch_oplog/metadata.cpp @@ -0,0 +1,409 @@ +#include "ha/snapshot/batch_oplog/metadata.h" + +#include +#include +#include +#include + +#if __has_include() +#include +#else +#include +#endif + +#include "ha/oplog/oplog_types.h" + +namespace mooncake::ha { + +namespace { + +void SetReason(std::string* reason, const std::string& value) { + if (reason != nullptr) { + *reason = value; + } +} + +std::string WriteJson(const Json::Value& root) { + Json::StreamWriterBuilder builder; + builder["indentation"] = ""; + return Json::writeString(builder, root); +} + +bool ParseJson(std::string_view value, Json::Value& root, std::string* reason) { + Json::CharReaderBuilder builder; + builder["allowComments"] = false; + builder["collectComments"] = false; + builder["failIfExtra"] = true; + builder["rejectDupKeys"] = true; + builder["strictRoot"] = true; + std::unique_ptr reader(builder.newCharReader()); + std::string errors; + if (!reader->parse(value.data(), value.data() + value.size(), &root, + &errors)) { + SetReason(reason, errors.empty() ? "malformed json" : errors); + return false; + } + if (!root.isObject()) { + SetReason(reason, "snapshot metadata must be a JSON object"); + return false; + } + return true; +} + +bool GetString(const Json::Value& root, const char* field, std::string& value, + std::string* reason) { + if (!root.isMember(field) || !root[field].isString()) { + SetReason(reason, std::string("field must be a string: ") + field); + return false; + } + value = root[field].asString(); + return true; +} + +bool GetUInt32(const Json::Value& root, const char* field, uint32_t& value, + std::string* reason) { + if (!root.isMember(field) || !root[field].isUInt()) { + SetReason(reason, std::string("field must be uint32: ") + field); + return false; + } + value = root[field].asUInt(); + return true; +} + +bool GetUInt64(const Json::Value& root, const char* field, uint64_t& value, + std::string* reason) { + if (!root.isMember(field) || !root[field].isUInt64()) { + SetReason(reason, std::string("field must be uint64: ") + field); + return false; + } + value = root[field].asUInt64(); + return true; +} + +template +bool GetInt64(const Json::Value& root, const char* field, Integer& value, + std::string* reason) { + if (!root.isMember(field) || !root[field].isInt64()) { + SetReason(reason, std::string("field must be int64: ") + field); + return false; + } + value = static_cast(root[field].asInt64()); + return true; +} + +template +bool ParseDecimal(std::string_view value, Integer& result) { + const char* begin = value.data(); + const char* end = begin + value.size(); + const auto parsed = std::from_chars(begin, end, result); + return begin != end && parsed.ec == std::errc() && parsed.ptr == end; +} + +bool ParseSnapshotId(std::string_view snapshot_id, uint64_t& batch_id) { + const size_t separator = snapshot_id.find('-'); + if (separator == std::string_view::npos || separator == 0 || + separator + 1 == snapshot_id.size() || + snapshot_id.find('-', separator + 1) != std::string_view::npos) { + return false; + } + + uint64_t parsed_batch_id = 0; + int64_t lease_id = 0; + if (!ParseDecimal(snapshot_id.substr(0, separator), parsed_batch_id) || + !ParseDecimal(snapshot_id.substr(separator + 1), lease_id) || + lease_id <= 0 || + snapshot_id != + std::to_string(parsed_batch_id) + "-" + std::to_string(lease_id)) { + return false; + } + batch_id = parsed_batch_id; + return true; +} + +template +bool ValidateIdentity(const SnapshotMetadata& metadata, std::string* reason) { + if (metadata.schema_version != kBatchOpLogSnapshotSchemaVersion) { + SetReason(reason, "unsupported snapshot schema_version"); + return false; + } + if (metadata.snapshot_format != kBatchOpLogSnapshotFormat) { + SetReason(reason, "unsupported snapshot_format"); + return false; + } + uint64_t snapshot_batch_id = 0; + if (!ParseSnapshotId(metadata.snapshot_id, snapshot_batch_id) || + snapshot_batch_id != metadata.last_included_batch_id) { + SetReason(reason, "snapshot_id does not match the batch cursor"); + return false; + } + if ((metadata.last_included_seq == 0) != + (metadata.last_included_batch_id == 0)) { + SetReason(reason, + "sequence and batch cursors must both be zero or non-zero"); + return false; + } + if (metadata.producer_view_version < 0) { + SetReason(reason, "producer_view_version must be non-negative"); + return false; + } + return true; +} + +template +void EncodeIdentity(const SnapshotMetadata& metadata, Json::Value& root) { + root["schema_version"] = static_cast(metadata.schema_version); + root["snapshot_format"] = metadata.snapshot_format; + root["snapshot_id"] = metadata.snapshot_id; + root["last_included_seq"] = + static_cast(metadata.last_included_seq); + root["last_included_batch_id"] = + static_cast(metadata.last_included_batch_id); + root["producer_view_version"] = + static_cast(metadata.producer_view_version); +} + +template +bool DecodeIdentity(const Json::Value& root, SnapshotMetadata& metadata, + std::string* reason) { + return GetUInt32(root, "schema_version", metadata.schema_version, reason) && + GetString(root, "snapshot_format", metadata.snapshot_format, + reason) && + GetString(root, "snapshot_id", metadata.snapshot_id, reason) && + GetUInt64(root, "last_included_seq", metadata.last_included_seq, + reason) && + GetUInt64(root, "last_included_batch_id", + metadata.last_included_batch_id, reason) && + GetInt64(root, "producer_view_version", + metadata.producer_view_version, reason) && + ValidateIdentity(metadata, reason); +} + +Json::Value EncodeObjectDescriptor( + const BatchOpLogSnapshotObjectDescriptor& descriptor) { + Json::Value root(Json::objectValue); + root["key"] = descriptor.key; + root["stored_size"] = static_cast(descriptor.stored_size); + root["crc32c"] = static_cast(descriptor.crc32c); + return root; +} + +bool DecodeObjectDescriptor(const Json::Value& root, + BatchOpLogSnapshotObjectDescriptor& descriptor, + std::string* reason) { + if (!root.isObject()) { + SetReason(reason, "segments must be a JSON object"); + return false; + } + if (!GetString(root, "key", descriptor.key, reason) || + !GetUInt64(root, "stored_size", descriptor.stored_size, reason) || + !GetUInt32(root, "crc32c", descriptor.crc32c, reason)) { + return false; + } + if (descriptor.key.empty() || descriptor.stored_size == 0) { + SetReason(reason, "segments key and stored_size must be non-zero"); + return false; + } + return true; +} + +Json::Value EncodeChunkDescriptor( + const BatchOpLogSnapshotChunkDescriptor& descriptor) { + Json::Value root(Json::objectValue); + root["chunk_index"] = static_cast(descriptor.chunk_index); + root["key"] = descriptor.key; + root["object_count"] = static_cast(descriptor.object_count); + root["stored_size"] = static_cast(descriptor.stored_size); + root["crc32c"] = static_cast(descriptor.crc32c); + return root; +} + +bool DecodeChunkDescriptor(const Json::Value& root, uint64_t expected_index, + BatchOpLogSnapshotChunkDescriptor& descriptor, + std::string* reason) { + if (!root.isObject()) { + SetReason(reason, "object chunk must be a JSON object"); + return false; + } + if (!GetUInt64(root, "chunk_index", descriptor.chunk_index, reason) || + !GetString(root, "key", descriptor.key, reason) || + !GetUInt64(root, "object_count", descriptor.object_count, reason) || + !GetUInt64(root, "stored_size", descriptor.stored_size, reason) || + !GetUInt32(root, "crc32c", descriptor.crc32c, reason)) { + return false; + } + if (descriptor.chunk_index != expected_index) { + SetReason(reason, "object chunk indices must be contiguous from zero"); + return false; + } + if (descriptor.key.empty() || descriptor.object_count == 0 || + descriptor.stored_size == 0) { + SetReason( + reason, + "object chunk key, object_count, and stored_size must be non-zero"); + return false; + } + return true; +} + +std::string BuildControlKey(const std::string& cluster_id, + std::string_view name) { + std::string normalized = cluster_id; + if (!NormalizeAndValidateClusterId(normalized) || normalized.empty()) { + return {}; + } + return "/oplog/" + normalized + "/snapshot/" + std::string(name); +} + +std::string BuildArtifactPrefix(const std::string& snapshot_root, + std::string_view snapshot_id) { + uint64_t ignored_batch_id = 0; + if (!ParseSnapshotId(snapshot_id, ignored_batch_id)) { + return {}; + } + std::string normalized = snapshot_root; + while (!normalized.empty() && normalized.back() == '/') { + normalized.pop_back(); + } + if (normalized.empty()) { + return {}; + } + return normalized + "/batch-oplog/" + std::string(snapshot_id) + "/"; +} + +} // namespace + +std::string EncodeBatchOpLogSnapshotDescriptor( + const BatchOpLogSnapshotDescriptor& descriptor) { + Json::Value root(Json::objectValue); + EncodeIdentity(descriptor, root); + root["manifest_key"] = descriptor.manifest_key; + root["manifest_size"] = static_cast(descriptor.manifest_size); + root["manifest_crc32c"] = + static_cast(descriptor.manifest_crc32c); + root["created_at_ms"] = static_cast(descriptor.created_at_ms); + return WriteJson(root); +} + +tl::expected +DecodeBatchOpLogSnapshotDescriptor(std::string_view value) { + std::string reason; + Json::Value root; + BatchOpLogSnapshotDescriptor decoded; + if (!ParseJson(value, root, &reason) || + !DecodeIdentity(root, decoded, &reason) || + !GetString(root, "manifest_key", decoded.manifest_key, &reason) || + !GetUInt64(root, "manifest_size", decoded.manifest_size, &reason) || + !GetUInt32(root, "manifest_crc32c", decoded.manifest_crc32c, &reason) || + !GetInt64(root, "created_at_ms", decoded.created_at_ms, &reason)) { + return tl::make_unexpected(std::move(reason)); + } + if (decoded.manifest_key.empty() || decoded.manifest_size == 0 || + decoded.created_at_ms < 0) { + SetReason(&reason, + "manifest key and size must be non-zero and created_at_ms " + "non-negative"); + return tl::make_unexpected(std::move(reason)); + } + return decoded; +} + +std::string EncodeBatchOpLogSnapshotManifest( + const BatchOpLogSnapshotManifest& manifest) { + Json::Value root(Json::objectValue); + EncodeIdentity(manifest, root); + root["segments"] = EncodeObjectDescriptor(manifest.segments); + Json::Value chunks(Json::arrayValue); + for (const auto& descriptor : manifest.object_chunks) { + chunks.append(EncodeChunkDescriptor(descriptor)); + } + root["object_chunks"] = std::move(chunks); + return WriteJson(root); +} + +tl::expected +DecodeBatchOpLogSnapshotManifest(std::string_view value) { + std::string reason; + Json::Value root; + BatchOpLogSnapshotManifest decoded; + if (!ParseJson(value, root, &reason) || + !DecodeIdentity(root, decoded, &reason)) { + return tl::make_unexpected(std::move(reason)); + } + if (!root.isMember("segments")) { + SetReason(&reason, "missing field: segments"); + return tl::make_unexpected(std::move(reason)); + } + if (!DecodeObjectDescriptor(root["segments"], decoded.segments, &reason)) { + return tl::make_unexpected(std::move(reason)); + } + if (!root.isMember("object_chunks") || !root["object_chunks"].isArray()) { + SetReason(&reason, "object_chunks must be a JSON array"); + return tl::make_unexpected(std::move(reason)); + } + const auto& chunks = root["object_chunks"]; + decoded.object_chunks.reserve(chunks.size()); + for (Json::ArrayIndex index = 0; index < chunks.size(); ++index) { + BatchOpLogSnapshotChunkDescriptor descriptor; + if (!DecodeChunkDescriptor(chunks[index], index, descriptor, &reason)) { + return tl::make_unexpected(std::move(reason)); + } + decoded.object_chunks.push_back(std::move(descriptor)); + } + return decoded; +} + +std::string BuildBatchOpLogSnapshotId(uint64_t last_included_batch_id, + int64_t maintenance_lease_id) { + if (maintenance_lease_id <= 0) { + return {}; + } + return std::to_string(last_included_batch_id) + "-" + + std::to_string(maintenance_lease_id); +} + +std::string BuildBatchOpLogSnapshotMaintenanceKey( + const std::string& cluster_id) { + return BuildControlKey(cluster_id, "maintenance"); +} + +std::string BuildBatchOpLogSnapshotLatestKey(const std::string& cluster_id) { + return BuildControlKey(cluster_id, "latest"); +} + +std::string BuildBatchOpLogSnapshotFallbackKey(const std::string& cluster_id) { + return BuildControlKey(cluster_id, "fallback"); +} + +std::string BuildBatchOpLogSnapshotCompactionFloorKey( + const std::string& cluster_id) { + return BuildControlKey(cluster_id, "compaction_floor"); +} + +std::string BuildBatchOpLogSnapshotDescriptorKey( + const std::string& snapshot_root, std::string_view snapshot_id) { + const std::string prefix = BuildArtifactPrefix(snapshot_root, snapshot_id); + return prefix.empty() ? std::string() : prefix + "descriptor.json"; +} + +std::string BuildBatchOpLogSnapshotManifestKey(const std::string& snapshot_root, + std::string_view snapshot_id) { + const std::string prefix = BuildArtifactPrefix(snapshot_root, snapshot_id); + return prefix.empty() ? std::string() : prefix + "manifest.json"; +} + +std::string BuildBatchOpLogSnapshotSegmentsKey(const std::string& snapshot_root, + std::string_view snapshot_id) { + const std::string prefix = BuildArtifactPrefix(snapshot_root, snapshot_id); + return prefix.empty() ? std::string() : prefix + "segments.bin"; +} + +std::string BuildBatchOpLogSnapshotObjectChunkKey( + const std::string& snapshot_root, std::string_view snapshot_id, + uint64_t chunk_index) { + const std::string prefix = BuildArtifactPrefix(snapshot_root, snapshot_id); + return prefix.empty() + ? std::string() + : prefix + "objects/" + std::to_string(chunk_index) + ".bin"; +} + +} // namespace mooncake::ha diff --git a/mooncake-store/tests/CMakeLists.txt b/mooncake-store/tests/CMakeLists.txt index 72e18be31f..56b6485ad5 100644 --- a/mooncake-store/tests/CMakeLists.txt +++ b/mooncake-store/tests/CMakeLists.txt @@ -156,6 +156,8 @@ add_store_test(snapshot_child_process_test ha/snapshot/snapshot_child_process_test.cpp) add_store_test(master_snapshot_codec_test ha/snapshot/master_snapshot_codec_test.cpp) +add_ha_test(batch_oplog_snapshot_types_test + ha/snapshot/batch_oplog/metadata_test.cpp) add_store_test(master_service_test_for_snapshot ha/snapshot/master_service_test_for_snapshot.cpp) add_store_test(non_ha_reconnect_test non_ha_reconnect_test.cpp) diff --git a/mooncake-store/tests/ha/snapshot/batch_oplog/metadata_test.cpp b/mooncake-store/tests/ha/snapshot/batch_oplog/metadata_test.cpp new file mode 100644 index 0000000000..9f669363a5 --- /dev/null +++ b/mooncake-store/tests/ha/snapshot/batch_oplog/metadata_test.cpp @@ -0,0 +1,267 @@ +#include "ha/snapshot/batch_oplog/metadata.h" + +#include + +#include +#include +#include + +#if __has_include() +#include +#else +#include +#endif + +namespace mooncake::ha::test { + +namespace { + +BatchOpLogSnapshotDescriptor MakeDescriptor() { + return { + .snapshot_id = "9-12345", + .last_included_seq = 42, + .last_included_batch_id = 9, + .producer_view_version = 7, + .manifest_key = "snapshots/batch-oplog/9-12345/manifest.json", + .manifest_size = 128, + .manifest_crc32c = 17, + .created_at_ms = 1700000000000, + }; +} + +BatchOpLogSnapshotManifest MakeManifest() { + return { + .snapshot_id = "9-12345", + .last_included_seq = 42, + .last_included_batch_id = 9, + .producer_view_version = 7, + .segments = {.key = "snapshots/batch-oplog/9-12345/segments.bin", + .stored_size = 64, + .crc32c = 18}, + .object_chunks = + { + {.chunk_index = 0, + .key = "snapshots/batch-oplog/9-12345/objects/0.bin", + .object_count = 100, + .stored_size = 1024, + .crc32c = 19}, + {.chunk_index = 1, + .key = "snapshots/batch-oplog/9-12345/objects/1.bin", + .object_count = 50, + .stored_size = 512, + .crc32c = 20}, + }, + }; +} + +Json::Value ParseJson(const std::string& value) { + Json::CharReaderBuilder builder; + Json::Value root; + std::string errors; + std::istringstream stream(value); + EXPECT_TRUE(Json::parseFromStream(builder, stream, &root, &errors)) + << errors; + return root; +} + +std::string WriteJson(const Json::Value& root) { + Json::StreamWriterBuilder builder; + builder["indentation"] = ""; + return Json::writeString(builder, root); +} + +void ExpectDescriptorRejected(const std::string& json) { + auto decoded = DecodeBatchOpLogSnapshotDescriptor(json); + ASSERT_FALSE(decoded.has_value()); + EXPECT_FALSE(decoded.error().empty()); +} + +void ExpectManifestRejected(const std::string& json) { + SCOPED_TRACE(json); + auto decoded = DecodeBatchOpLogSnapshotManifest(json); + ASSERT_FALSE(decoded.has_value()); + EXPECT_FALSE(decoded.error().empty()); +} + +} // namespace + +TEST(BatchOpLogSnapshotTypesTest, DescriptorRoundTripsCompactJson) { + const auto encoded = EncodeBatchOpLogSnapshotDescriptor(MakeDescriptor()); + EXPECT_EQ(encoded.find('\n'), std::string::npos); + + auto json = ParseJson(encoded); + json["future_optional_field"] = true; + + auto decoded = DecodeBatchOpLogSnapshotDescriptor(WriteJson(json)); + ASSERT_TRUE(decoded.has_value()) << decoded.error(); + EXPECT_EQ(decoded->schema_version, kBatchOpLogSnapshotSchemaVersion); + EXPECT_EQ(decoded->snapshot_format, kBatchOpLogSnapshotFormat); + EXPECT_EQ(decoded->snapshot_id, "9-12345"); + EXPECT_EQ(decoded->last_included_seq, 42u); + EXPECT_EQ(decoded->last_included_batch_id, 9u); + EXPECT_EQ(decoded->producer_view_version, 7u); + EXPECT_EQ(decoded->manifest_key, + "snapshots/batch-oplog/9-12345/manifest.json"); + EXPECT_EQ(decoded->manifest_size, 128u); + EXPECT_EQ(decoded->manifest_crc32c, 17u); + EXPECT_EQ(decoded->created_at_ms, 1700000000000); +} + +TEST(BatchOpLogSnapshotTypesTest, ManifestRoundTripsChunksAndAllowsEmptySet) { + auto encoded = EncodeBatchOpLogSnapshotManifest(MakeManifest()); + EXPECT_EQ(encoded.find('\n'), std::string::npos); + auto json = ParseJson(encoded); + json["future_optional_field"] = true; + auto decoded = DecodeBatchOpLogSnapshotManifest(WriteJson(json)); + ASSERT_TRUE(decoded.has_value()) << decoded.error(); + ASSERT_EQ(decoded->object_chunks.size(), 2u); + EXPECT_EQ(decoded->segments.stored_size, 64u); + EXPECT_EQ(decoded->object_chunks[1].chunk_index, 1u); + EXPECT_EQ(decoded->object_chunks[1].object_count, 50u); + + auto empty = MakeManifest(); + empty.object_chunks.clear(); + auto decoded_empty = DecodeBatchOpLogSnapshotManifest( + EncodeBatchOpLogSnapshotManifest(empty)); + ASSERT_TRUE(decoded_empty.has_value()) << decoded_empty.error(); + EXPECT_TRUE(decoded_empty->object_chunks.empty()); +} + +TEST(BatchOpLogSnapshotTypesTest, JsonRoundTripPreservesEscapedKeys) { + auto descriptor = MakeDescriptor(); + descriptor.manifest_key = R"(snapshots/"quoted"/manifest\key.json)"; + auto decoded_descriptor = DecodeBatchOpLogSnapshotDescriptor( + EncodeBatchOpLogSnapshotDescriptor(descriptor)); + ASSERT_TRUE(decoded_descriptor.has_value()) << decoded_descriptor.error(); + EXPECT_EQ(decoded_descriptor->manifest_key, descriptor.manifest_key); + + auto manifest = MakeManifest(); + manifest.segments.key = R"(snapshots/"quoted"/segments\key.bin)"; + manifest.object_chunks[0].key = R"(snapshots/"quoted"/objects\0.bin)"; + auto decoded_manifest = DecodeBatchOpLogSnapshotManifest( + EncodeBatchOpLogSnapshotManifest(manifest)); + ASSERT_TRUE(decoded_manifest.has_value()) << decoded_manifest.error(); + EXPECT_EQ(decoded_manifest->segments.key, manifest.segments.key); + EXPECT_EQ(decoded_manifest->object_chunks[0].key, + manifest.object_chunks[0].key); +} + +TEST(BatchOpLogSnapshotTypesTest, RejectsInvalidDescriptorJson) { + const auto encoded = EncodeBatchOpLogSnapshotDescriptor(MakeDescriptor()); + auto json = ParseJson(encoded); + + ExpectDescriptorRejected("{"); + ExpectDescriptorRejected("[]"); + ExpectDescriptorRejected(encoded.substr(0, encoded.size() - 1) + + ",\"schema_version\":1}"); + + auto invalid = json; + invalid.removeMember("manifest_key"); + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["manifest_size"] = "128"; + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["schema_version"] = 2; + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["snapshot_format"] = "standby-oplog-materialized/v2"; + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["last_included_batch_id"] = Json::UInt64(0); + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["snapshot_id"] = "10-12345"; + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["manifest_key"] = ""; + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["manifest_size"] = Json::UInt64(0); + ExpectDescriptorRejected(WriteJson(invalid)); + invalid = json; + invalid["created_at_ms"] = Json::Int64(-1); + ExpectDescriptorRejected(WriteJson(invalid)); + + auto overflow = encoded; + const std::string old_value = "\"last_included_seq\":42"; + const auto position = overflow.find(old_value); + ASSERT_NE(position, std::string::npos); + overflow.replace(position, old_value.size(), + "\"last_included_seq\":18446744073709551616"); + ExpectDescriptorRejected(overflow); +} + +TEST(BatchOpLogSnapshotTypesTest, AcceptsAnEmptyCursor) { + auto descriptor = MakeDescriptor(); + descriptor.snapshot_id = "0-12345"; + descriptor.last_included_seq = 0; + descriptor.last_included_batch_id = 0; + + auto decoded = DecodeBatchOpLogSnapshotDescriptor( + EncodeBatchOpLogSnapshotDescriptor(descriptor)); + ASSERT_TRUE(decoded.has_value()) << decoded.error(); +} + +TEST(BatchOpLogSnapshotTypesTest, RejectsInvalidManifestJson) { + const auto encoded = EncodeBatchOpLogSnapshotManifest(MakeManifest()); + auto json = ParseJson(encoded); + + ExpectManifestRejected(encoded.substr(0, encoded.size() - 1) + + ",\"schema_version\":1}"); + + auto invalid = json; + invalid.removeMember("segments"); + ExpectManifestRejected(WriteJson(invalid)); + invalid = json; + invalid["segments"]["key"] = ""; + ExpectManifestRejected(WriteJson(invalid)); + invalid = json; + invalid["segments"]["crc32c"] = + Json::UInt64(static_cast(UINT32_MAX) + 1); + ExpectManifestRejected(WriteJson(invalid)); + invalid = json; + invalid["object_chunks"][1]["chunk_index"] = Json::UInt64(2); + ExpectManifestRejected(WriteJson(invalid)); + invalid = json; + invalid["object_chunks"][0]["object_count"] = Json::UInt64(0); + ExpectManifestRejected(WriteJson(invalid)); + invalid = json; + invalid["snapshot_id"] = "8-12345"; + ExpectManifestRejected(WriteJson(invalid)); +} + +TEST(BatchOpLogSnapshotTypesTest, BuildsControlAndArtifactKeys) { + EXPECT_EQ(BuildBatchOpLogSnapshotMaintenanceKey("cluster-a/"), + "/oplog/cluster-a/snapshot/maintenance"); + EXPECT_EQ(BuildBatchOpLogSnapshotLatestKey("cluster-a"), + "/oplog/cluster-a/snapshot/latest"); + EXPECT_EQ(BuildBatchOpLogSnapshotFallbackKey("cluster-a"), + "/oplog/cluster-a/snapshot/fallback"); + EXPECT_EQ(BuildBatchOpLogSnapshotCompactionFloorKey("cluster-a"), + "/oplog/cluster-a/snapshot/compaction_floor"); + EXPECT_TRUE(BuildBatchOpLogSnapshotLatestKey("../cluster").empty()); + + const auto snapshot_id = BuildBatchOpLogSnapshotId(9, 12345); + ASSERT_EQ(snapshot_id, "9-12345"); + EXPECT_EQ(BuildBatchOpLogSnapshotDescriptorKey("snapshots", snapshot_id), + "snapshots/batch-oplog/9-12345/descriptor.json"); + EXPECT_EQ(BuildBatchOpLogSnapshotManifestKey("snapshots/", snapshot_id), + "snapshots/batch-oplog/9-12345/manifest.json"); + EXPECT_EQ(BuildBatchOpLogSnapshotSegmentsKey("snapshots", snapshot_id), + "snapshots/batch-oplog/9-12345/segments.bin"); + EXPECT_EQ( + BuildBatchOpLogSnapshotObjectChunkKey("snapshots", snapshot_id, 7), + "snapshots/batch-oplog/9-12345/objects/7.bin"); +} + +TEST(BatchOpLogSnapshotTypesTest, RejectsUnsafeSnapshotIdInArtifactKeys) { + EXPECT_TRUE(BuildBatchOpLogSnapshotId(9, 0).empty()); + EXPECT_TRUE( + BuildBatchOpLogSnapshotDescriptorKey("snapshots", "9-x").empty()); + EXPECT_TRUE(BuildBatchOpLogSnapshotDescriptorKey("snapshots", + "9-12345/../../latest") + .empty()); +} + +} // namespace mooncake::ha::test From 9187a2cf9a3fd4f3c3c41f640a8b28c086e50ccc Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Tue, 11 Aug 2026 11:41:45 +0800 Subject: [PATCH 025/483] [Store] Add producer view to durable oplog prefix (#3201) Co-authored-by: Yuchen Kou --- .../include/ha/oplog/oplog_batch_types.h | 6 ++ .../src/ha/oplog/oplog_batch_codec.cpp | 20 +++- .../src/ha/oplog/oplog_batch_storage.cpp | 31 ++++-- .../src/ha/oplog/ordered_oplog_writer.cpp | 7 +- .../tests/ha/oplog/oplog_batch_codec_test.cpp | 94 +++++++++++++++++++ .../ha/oplog/oplog_batch_storage_test.cpp | 87 +++++++++++++++++ .../ha/oplog/ordered_oplog_writer_test.cpp | 43 +++++++++ 7 files changed, 276 insertions(+), 12 deletions(-) diff --git a/mooncake-store/include/ha/oplog/oplog_batch_types.h b/mooncake-store/include/ha/oplog/oplog_batch_types.h index 4e315383c1..a5442176bd 100644 --- a/mooncake-store/include/ha/oplog/oplog_batch_types.h +++ b/mooncake-store/include/ha/oplog/oplog_batch_types.h @@ -15,6 +15,12 @@ static constexpr int kOpLogBatchIdWidth = 20; struct DurablePrefix { uint64_t batch_id{0}; uint64_t last_seq{0}; + // Leadership view that produced this durable boundary. Zero means that + // the prefix has no producer-view metadata (legacy or not yet assigned). + ViewVersionId producer_view_version{0}; + + friend bool operator==(const DurablePrefix&, + const DurablePrefix&) = default; }; struct BatchRecordRange { diff --git a/mooncake-store/src/ha/oplog/oplog_batch_codec.cpp b/mooncake-store/src/ha/oplog/oplog_batch_codec.cpp index 5c5a644bfa..aae0b194c8 100644 --- a/mooncake-store/src/ha/oplog/oplog_batch_codec.cpp +++ b/mooncake-store/src/ha/oplog/oplog_batch_codec.cpp @@ -153,6 +153,10 @@ std::string EncodeDurablePrefix(const DurablePrefix& prefix) { static_cast(kDurablePrefixSchemaVersion); root["batch_id"] = static_cast(prefix.batch_id); root["last_seq"] = static_cast(prefix.last_seq); + if (prefix.producer_view_version != 0) { + root["producer_view_version"] = + static_cast(prefix.producer_view_version); + } return WriteJson(root); } @@ -182,12 +186,24 @@ bool DecodeDurablePrefix(const std::string& value, DurablePrefix* prefix, SetReason(reason, "unsupported durable prefix schema_version"); return false; } - if (!GetUInt64Field(root, "batch_id", &prefix->batch_id, reason)) { + DurablePrefix decoded; + if (root.isMember("producer_view_version")) { + const auto& producer_view = root["producer_view_version"]; + if (!producer_view.isInt64() || producer_view.asInt64() < 0) { + SetReason(reason, + "field must be a non-negative ViewVersionId: " + "producer_view_version"); + return false; + } + decoded.producer_view_version = producer_view.asInt64(); + } + if (!GetUInt64Field(root, "batch_id", &decoded.batch_id, reason)) { return false; } - if (!GetUInt64Field(root, "last_seq", &prefix->last_seq, reason)) { + if (!GetUInt64Field(root, "last_seq", &decoded.last_seq, reason)) { return false; } + *prefix = decoded; return true; } diff --git a/mooncake-store/src/ha/oplog/oplog_batch_storage.cpp b/mooncake-store/src/ha/oplog/oplog_batch_storage.cpp index 8229b2262e..5c74ea3fd2 100644 --- a/mooncake-store/src/ha/oplog/oplog_batch_storage.cpp +++ b/mooncake-store/src/ha/oplog/oplog_batch_storage.cpp @@ -104,8 +104,7 @@ ErrorCode OpLogBatchStorage::InitDurablePrefix(DurablePrefix& prefix) { return ErrorCode::OK; } if (err == ErrorCode::ETCD_TRANSACTION_FAIL) { - err = ReadDurablePrefix(prefix); - if (err != ErrorCode::OK) { + if ((err = ReadDurablePrefix(prefix)) != ErrorCode::OK) { return err; } return ValidateDurablePrefixAtStartup(prefix); @@ -219,19 +218,35 @@ ErrorCode OpLogBatchStorage::WriteBatchAndAdvancePrefix( .expected_value = EncodeDurablePrefix(expected_prefix)}); txn.puts.push_back({.key = BuildBatchRecordKey(cluster_id_, batch.batch_id), .value = encoded_batch}); - txn.puts.push_back( - {.key = durable_key, - .value = EncodeDurablePrefix( - {.batch_id = batch.batch_id, .last_seq = batch.last_seq})}); + txn.puts.push_back({.key = durable_key, + .value = EncodeDurablePrefix( + {.batch_id = batch.batch_id, + .last_seq = batch.last_seq, + .producer_view_version = + expected_prefix.producer_view_version})}); ErrorCode err = backend_.Txn(txn); + if (err == ErrorCode::ETCD_TRANSACTION_FAIL) { + std::string raw_prefix; + DurablePrefix decoded_prefix; + if (backend_.Get(durable_key, raw_prefix) == ErrorCode::OK && + raw_prefix != txn.compares[0].expected_value && + DecodeDurablePrefix(raw_prefix, &decoded_prefix) && + decoded_prefix == expected_prefix) { + txn.compares[0].expected_value = raw_prefix; + err = backend_.Txn(txn); + } + } if (err != ErrorCode::ETCD_TRANSACTION_FAIL) { return err; } DurablePrefix current_prefix; if (ReadDurablePrefix(current_prefix) != ErrorCode::OK || - current_prefix.batch_id != batch.batch_id || - current_prefix.last_seq != batch.last_seq) { + current_prefix != + DurablePrefix{.batch_id = batch.batch_id, + .last_seq = batch.last_seq, + .producer_view_version = + expected_prefix.producer_view_version}) { return err; } OpLogBatchRecord current_batch; diff --git a/mooncake-store/src/ha/oplog/ordered_oplog_writer.cpp b/mooncake-store/src/ha/oplog/ordered_oplog_writer.cpp index 43807a4703..f1953bb44f 100644 --- a/mooncake-store/src/ha/oplog/ordered_oplog_writer.cpp +++ b/mooncake-store/src/ha/oplog/ordered_oplog_writer.cpp @@ -319,8 +319,11 @@ void OrderedOpLogWriter::Start() { #endif { std::lock_guard lock(impl_->mutex); - impl_->durable_prefix = {.batch_id = batch.batch_id, - .last_seq = batch.last_seq}; + impl_->durable_prefix = { + .batch_id = batch.batch_id, + .last_seq = batch.last_seq, + .producer_view_version = + expected_prefix.producer_view_version}; impl_->last_error = ErrorCode::OK; impl_->accepting = !impl_->stop_requested; for (size_t i = 0; i < entries.size(); ++i) { diff --git a/mooncake-store/tests/ha/oplog/oplog_batch_codec_test.cpp b/mooncake-store/tests/ha/oplog/oplog_batch_codec_test.cpp index 22a1bf8b64..18d10c3ad2 100644 --- a/mooncake-store/tests/ha/oplog/oplog_batch_codec_test.cpp +++ b/mooncake-store/tests/ha/oplog/oplog_batch_codec_test.cpp @@ -4,6 +4,7 @@ #include #include +#include #include #include #include @@ -178,6 +179,99 @@ TEST(OpLogDurablePrefixCodecTest, RoundTripsNonZeroPrefix) { EXPECT_TRUE(reason.empty()); } +TEST(OpLogDurablePrefixCodecTest, PreservesLegacyEncodingWithoutProducerView) { + const std::string legacy = + R"({"batch_id":9,"last_seq":1024,"schema_version":1})"; + + DurablePrefix prefix; + std::string reason; + ASSERT_TRUE(DecodeDurablePrefix(legacy, &prefix, &reason)); + EXPECT_EQ(0u, prefix.producer_view_version); + EXPECT_EQ(legacy, EncodeDurablePrefix(prefix)); + EXPECT_TRUE(reason.empty()); +} + +TEST(OpLogDurablePrefixCodecTest, RoundTripsMaxProducerView) { + DurablePrefix in{ + .batch_id = 9, + .last_seq = 1024, + .producer_view_version = std::numeric_limits::max()}; + + DurablePrefix out; + std::string reason; + ASSERT_TRUE(DecodeDurablePrefix(EncodeDurablePrefix(in), &out, &reason)); + EXPECT_EQ(in.producer_view_version, out.producer_view_version); + EXPECT_TRUE(reason.empty()); +} + +TEST(OpLogDurablePrefixCodecTest, OmitsExplicitZeroProducerViewOnReencode) { + const std::string with_explicit_zero = + R"({"batch_id":9,"last_seq":1024,"producer_view_version":0,"schema_version":1})"; + const std::string without_view = + R"({"batch_id":9,"last_seq":1024,"schema_version":1})"; + + DurablePrefix prefix; + std::string reason; + ASSERT_TRUE(DecodeDurablePrefix(with_explicit_zero, &prefix, &reason)); + EXPECT_EQ(0u, prefix.producer_view_version); + EXPECT_EQ(without_view, EncodeDurablePrefix(prefix)); + EXPECT_TRUE(reason.empty()); +} + +TEST(OpLogDurablePrefixCodecTest, IgnoresUnknownFields) { + DurablePrefix prefix; + std::string reason; + ASSERT_TRUE(DecodeDurablePrefix( + R"({"batch_id":9,"last_seq":1024,"producer_view_version":7,"schema_version":1,"unknown":"ignored"})", + &prefix, &reason)); + EXPECT_EQ(9u, prefix.batch_id); + EXPECT_EQ(1024u, prefix.last_seq); + EXPECT_EQ(7u, prefix.producer_view_version); + EXPECT_TRUE(reason.empty()); +} + +TEST(OpLogDurablePrefixCodecTest, RejectsInvalidProducerViews) { + const std::vector invalid_values = { + R"("1")", "-1", "1.5", "9223372036854775808", "18446744073709551616"}; + + for (const auto& invalid_value : invalid_values) { + DurablePrefix prefix; + std::string reason; + const std::string encoded = + R"({"batch_id":9,"last_seq":1024,"producer_view_version":)" + + invalid_value + R"(,"schema_version":1})"; + EXPECT_FALSE(DecodeDurablePrefix(encoded, &prefix, &reason)) + << invalid_value; + EXPECT_FALSE(reason.empty()) << invalid_value; + } +} + +TEST(OpLogDurablePrefixCodecTest, InvalidProducerViewLeavesOutputUnchanged) { + DurablePrefix prefix{ + .batch_id = 11, .last_seq = 22, .producer_view_version = 33}; + std::string reason; + EXPECT_FALSE(DecodeDurablePrefix( + R"({"batch_id":9,"last_seq":1024,"producer_view_version":"invalid","schema_version":1})", + &prefix, &reason)); + EXPECT_EQ(11u, prefix.batch_id); + EXPECT_EQ(22u, prefix.last_seq); + EXPECT_EQ(33u, prefix.producer_view_version); + EXPECT_FALSE(reason.empty()); +} + +TEST(OpLogDurablePrefixCodecTest, InvalidLastSeqLeavesOutputUnchanged) { + DurablePrefix prefix{ + .batch_id = 11, .last_seq = 22, .producer_view_version = 33}; + std::string reason; + EXPECT_FALSE(DecodeDurablePrefix( + R"({"batch_id":9,"last_seq":"invalid","schema_version":1})", &prefix, + &reason)); + EXPECT_EQ(11u, prefix.batch_id); + EXPECT_EQ(22u, prefix.last_seq); + EXPECT_EQ(33u, prefix.producer_view_version); + EXPECT_FALSE(reason.empty()); +} + TEST(OpLogDurablePrefixCodecTest, RejectsMalformedPayload) { DurablePrefix out; std::string reason; diff --git a/mooncake-store/tests/ha/oplog/oplog_batch_storage_test.cpp b/mooncake-store/tests/ha/oplog/oplog_batch_storage_test.cpp index c1df73d21a..2a9742b8d0 100644 --- a/mooncake-store/tests/ha/oplog/oplog_batch_storage_test.cpp +++ b/mooncake-store/tests/ha/oplog/oplog_batch_storage_test.cpp @@ -169,6 +169,20 @@ TEST(OpLogBatchStorageTest, InitializesEmptyNamespaceAtZero) { EXPECT_EQ(prefix.last_seq, stored.last_seq); } +TEST(OpLogBatchStorageTest, AcceptsNonzeroViewAtEmptyBoundary) { + FakeHaKvBackend backend; + ASSERT_EQ(ErrorCode::OK, + backend.Put("/oplog/clusterA/durable_prefix", + EncodeDurablePrefix({.batch_id = 0, + .last_seq = 0, + .producer_view_version = 7}))); + OpLogBatchStorage storage("clusterA", backend); + + DurablePrefix prefix; + EXPECT_EQ(ErrorCode::OK, storage.InitDurablePrefix(prefix)); + EXPECT_EQ(7u, prefix.producer_view_version); +} + TEST(OpLogBatchStorageTest, RejectsLegacyLatest) { FakeHaKvBackend backend; ASSERT_EQ(ErrorCode::OK, backend.Put("/oplog/clusterA/latest", "42")); @@ -379,6 +393,58 @@ TEST(OpLogBatchStorageTest, WriteBatchAndAdvancePrefixCommitsAtomically) { EXPECT_EQ(5u, prefix.last_seq); } +TEST(OpLogBatchStorageTest, PreservesProducerViewWhenAdvancingPrefix) { + FakeHaKvBackend backend; + const DurablePrefix expected_prefix{ + .batch_id = 1, .last_seq = 3, .producer_view_version = 7}; + ASSERT_EQ(ErrorCode::OK, backend.Put("/oplog/clusterA/durable_prefix", + EncodeDurablePrefix(expected_prefix))); + OpLogBatchStorage storage("clusterA", backend); + + auto batch = MakeBatch(/*batch_id=*/2, /*first_seq=*/4, /*count=*/2); + ASSERT_EQ(ErrorCode::OK, + storage.WriteBatchAndAdvancePrefix(batch, expected_prefix)); + + std::string encoded_prefix; + ASSERT_EQ(ErrorCode::OK, + backend.Get("/oplog/clusterA/durable_prefix", encoded_prefix)); + DurablePrefix prefix; + ASSERT_TRUE(DecodeDurablePrefix(encoded_prefix, &prefix)); + EXPECT_EQ(2u, prefix.batch_id); + EXPECT_EQ(5u, prefix.last_seq); + EXPECT_EQ(7u, prefix.producer_view_version); +} + +TEST(OpLogBatchStorageTest, ExplicitZeroProducerViewCanAdvancePrefix) { + FakeHaKvBackend backend; + ASSERT_EQ(ErrorCode::OK, + backend.Put("/oplog/clusterA/batches/00000000000000000001", + EncodeOpLogBatchRecord(MakeBatch( + /*batch_id=*/1, /*first_seq=*/1, /*count=*/3)))); + ASSERT_EQ( + ErrorCode::OK, + backend.Put( + "/oplog/clusterA/durable_prefix", + R"({"schema_version":1,"batch_id":1,"last_seq":3,"producer_view_version":0})")); + OpLogBatchStorage storage("clusterA", backend); + + DurablePrefix prefix; + ASSERT_EQ(ErrorCode::OK, storage.InitDurablePrefix(prefix)); + ASSERT_EQ(0u, prefix.producer_view_version); + ASSERT_EQ( + ErrorCode::OK, + storage.WriteBatchAndAdvancePrefix( + MakeBatch(/*batch_id=*/2, /*first_seq=*/4, /*count=*/2), prefix)); + + std::string encoded_prefix; + ASSERT_EQ(ErrorCode::OK, + backend.Get("/oplog/clusterA/durable_prefix", encoded_prefix)); + ASSERT_TRUE(DecodeDurablePrefix(encoded_prefix, &prefix)); + EXPECT_EQ(2u, prefix.batch_id); + EXPECT_EQ(5u, prefix.last_seq); + EXPECT_EQ(0u, prefix.producer_view_version); +} + TEST(OpLogBatchStorageTest, CompareFailureDoesNotWriteBatchOrAdvancePrefix) { FakeHaKvBackend backend; ASSERT_EQ(ErrorCode::OK, @@ -418,6 +484,27 @@ TEST(OpLogBatchStorageTest, CompareFailureIsOkWhenTargetBatchAlreadyDurable) { batch, {.batch_id = 1, .last_seq = 3})); } +TEST(OpLogBatchStorageTest, + CompareFailureWithDifferentProducerViewIsNotIdempotentSuccess) { + FakeHaKvBackend backend; + ASSERT_EQ(ErrorCode::OK, + backend.Put("/oplog/clusterA/durable_prefix", + EncodeDurablePrefix({.batch_id = 2, + .last_seq = 5, + .producer_view_version = 8}))); + auto batch = MakeBatch(/*batch_id=*/2, /*first_seq=*/4, /*count=*/2); + ASSERT_EQ(ErrorCode::OK, + backend.Put("/oplog/clusterA/batches/00000000000000000002", + EncodeOpLogBatchRecord(batch))); + backend.FailNextTxn(ErrorCode::ETCD_TRANSACTION_FAIL); + OpLogBatchStorage storage("clusterA", backend); + + EXPECT_EQ( + ErrorCode::ETCD_TRANSACTION_FAIL, + storage.WriteBatchAndAdvancePrefix( + batch, {.batch_id = 1, .last_seq = 3, .producer_view_version = 7})); +} + TEST(OpLogBatchStorageTest, RejectsSkippedBatchId) { FakeHaKvBackend backend; ASSERT_EQ(ErrorCode::OK, diff --git a/mooncake-store/tests/ha/oplog/ordered_oplog_writer_test.cpp b/mooncake-store/tests/ha/oplog/ordered_oplog_writer_test.cpp index 29c2c011a7..3b3412f619 100644 --- a/mooncake-store/tests/ha/oplog/ordered_oplog_writer_test.cpp +++ b/mooncake-store/tests/ha/oplog/ordered_oplog_writer_test.cpp @@ -57,6 +57,16 @@ class FakeBatchWriter { return batches; } + std::vector ExpectedPrefixes() const { + std::lock_guard lock(mutex_); + std::vector prefixes; + prefixes.reserve(writes_.size()); + for (const auto& write : writes_) { + prefixes.push_back(write.expected_prefix); + } + return prefixes; + } + bool WaitForWrites(size_t count, std::chrono::milliseconds timeout = std::chrono::milliseconds(1000)) { std::unique_lock lock(mutex_); @@ -589,6 +599,39 @@ TEST(OrderedOpLogWriterLoopTest, ContinuesFromInitialDurablePrefix) { writer.Stop(); } +TEST(OrderedOpLogWriterLoopTest, PreservesProducerViewAcrossBatches) { + FakeBatchWriter storage; + OrderedOpLogWriter writer( + OrderedOpLogWriterConfig{ + .max_entries_per_batch = 1, + .initial_durable_prefix = {.producer_view_version = 7}}, + [&](const OpLogBatchRecord& batch, + const DurablePrefix& expected_prefix) { + return storage.Write(batch, expected_prefix); + }); + writer.Start(); + + auto first = writer.Reserve(); + ASSERT_TRUE(first.has_value()); + ASSERT_TRUE( + writer.Commit(std::move(*first), MakeEntry("k1"), [](const auto&) {}) + .has_value()); + ASSERT_TRUE(storage.WaitForWrites(1)); + + auto second = writer.Reserve(); + ASSERT_TRUE(second.has_value()); + ASSERT_TRUE( + writer.Commit(std::move(*second), MakeEntry("k2"), [](const auto&) {}) + .has_value()); + ASSERT_TRUE(storage.WaitForWrites(2)); + + const auto prefixes = storage.ExpectedPrefixes(); + ASSERT_EQ(2u, prefixes.size()); + EXPECT_EQ(7u, prefixes[0].producer_view_version); + EXPECT_EQ(7u, prefixes[1].producer_view_version); + writer.Stop(); +} + TEST(OrderedOpLogWriterLoopTest, CommitWhileReadyBatchExistsFormsNextBatch) { FakeBatchWriter storage; OrderedOpLogWriter writer( From c6b3f632c6e81a1ac71ee73a14108f8e1414299a Mon Sep 17 00:00:00 2001 From: LZW <99333079+Lin-z-w@users.noreply.github.com> Date: Tue, 11 Aug 2026 11:48:37 +0800 Subject: [PATCH 026/483] [Store] Add fixed soft pin lifecycle (#2909) --- .../api-reference/cpp/mooncake-store.md | 12 +- .../api-reference/python/mooncake-store.md | 26 +- .../mooncake-store-deployment-guide.md | 1 + docs/source/design/mooncake-store.md | 22 +- docs/source/zh_archive/mooncake-store.md | 16 +- mooncake-integration/store/store_py.cpp | 8 +- .../store/store_py_internal.h | 6 +- mooncake-store/conf/master.json | 1 + mooncake-store/conf/master.yaml | 1 + mooncake-store/include/master_config.h | 15 + mooncake-store/include/master_service.h | 292 +++++++- mooncake-store/include/metadata_store.h | 2 + mooncake-store/include/replica.h | 37 +- mooncake-store/include/types.h | 2 + .../catalog_backed_snapshot_provider.cpp | 25 +- mooncake-store/src/master.cpp | 19 + mooncake-store/src/master_service.cpp | 446 +++++++++--- mooncake-store/src/store_c.cpp | 4 +- mooncake-store/tests/batch_evict_test.cpp | 3 +- .../catalog_backed_snapshot_provider_test.cpp | 26 +- .../master_service_test_for_snapshot.cpp | 122 +--- .../snapshot/snapshot_child_process_test.cpp | 53 +- .../tests/ha/snapshot/snapshot_test_utils.h | 19 +- mooncake-store/tests/master_service_test.cpp | 652 +++++++++++++++--- .../tests/test_distributed_object_store.py | 57 +- .../test_distributed_object_store_cxl.py | 8 +- mooncake-wheel/tests/test_dummy_client.py | 8 +- ...est_replicated_distributed_object_store.py | 6 +- 28 files changed, 1532 insertions(+), 357 deletions(-) diff --git a/docs/source/api-reference/cpp/mooncake-store.md b/docs/source/api-reference/cpp/mooncake-store.md index 779655148c..aca0e97953 100644 --- a/docs/source/api-reference/cpp/mooncake-store.md +++ b/docs/source/api-reference/cpp/mooncake-store.md @@ -48,12 +48,22 @@ The data structure details of `ReplicateConfig` are as follows: ```C++ struct ReplicateConfig { size_t replica_num{1}; // Total number of replicas for the object - bool with_soft_pin{false}; // Whether to enable soft pin mechanism for this object + SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; + std::optional soft_pin_ttl_ms{}; // ENABLE override; omitted uses the Master default bool with_hard_pin{false}; // Whether to enable hard pin (never evicted) std::string preferred_segment{}; // Preferred segment for allocation }; ``` +Soft pinning starts when the first replica becomes readable and has a fixed +lifetime: reads do not extend it. `PRESERVE` keeps the committed deadline on an +Upsert, `ENABLE` starts a new lifetime, and `DISABLE` removes it when the write +commits. `soft_pin_ttl_ms` is valid only with `ENABLE`; zero commits ordinary +cache, and values above the Master's configured maximum are rejected. +Soft-pin state is not persisted in snapshots or the HA OpLog; after recovery or +Standby promotion, restored objects are ordinary cache until a later write +explicitly enables soft pinning again. + ### Upsert ```C++ diff --git a/docs/source/api-reference/python/mooncake-store.md b/docs/source/api-reference/python/mooncake-store.md index 96094b015a..bbce3845e7 100644 --- a/docs/source/api-reference/python/mooncake-store.md +++ b/docs/source/api-reference/python/mooncake-store.md @@ -609,16 +609,24 @@ config = ReplicateConfig() config.replica_num = 3 # Store 3 copies of the data ``` -#### with_soft_pin -**Type:** `bool` -**Default:** `False` -**Description:** Enables soft pinning for the stored object. Soft pinned objects are prioritized to remain in memory during eviction - they are only evicted when memory is insufficient and no other objects are eligible for eviction. This is useful for frequently accessed or important objects like system prompts. +#### soft_pin_action +**Type:** `SoftPinAction` +**Default:** `SoftPinAction.PRESERVE` +**Description:** Controls the soft-pin transition committed when the first replica becomes readable. `PRESERVE` keeps an existing deadline during Upsert, `ENABLE` starts a fixed soft-pin lifetime, and `DISABLE` removes it. Reads do not extend the lifetime. ```python +from mooncake.store import ReplicateConfig, SoftPinAction + config = ReplicateConfig() -config.with_soft_pin = True # Keep this object in memory longer +config.soft_pin_action = SoftPinAction.ENABLE +config.soft_pin_ttl_ms = 60_000 # Optional; omitted uses the Master default ``` +`soft_pin_ttl_ms` is valid only with `ENABLE`. The Master rejects TTLs above +`max_kv_soft_pin_ttl`; a value of zero commits the object as ordinary cache. +Soft-pin state is not persisted in snapshots or the HA OpLog. Restored objects +therefore become ordinary cache after recovery or Standby promotion. + #### with_hard_pin **Type:** `bool` **Default:** `False` @@ -2161,7 +2169,7 @@ def pub_tensor(self, key: str, tensor: torch.Tensor, config: ReplicateConfig = N **Example:** ```python import torch -from mooncake.store import ReplicateConfig +from mooncake.store import ReplicateConfig, SoftPinAction # Create a tensor tensor = torch.randn(100, 100) @@ -2169,7 +2177,7 @@ tensor = torch.randn(100, 100) # Create replication config config = ReplicateConfig() config.replica_num = 3 -config.with_soft_pin = True +config.soft_pin_action = SoftPinAction.ENABLE # Publish tensor with replication settings result = store.pub_tensor("my_tensor", tensor, config) @@ -2575,13 +2583,13 @@ shared-memory staging buffer. **Example:** ```python import torch -from mooncake.store import ReplicateConfig +from mooncake.store import ReplicateConfig, SoftPinAction tensor = torch.randn(100, 100) config = ReplicateConfig() config.replica_num = 2 -config.with_soft_pin = True +config.soft_pin_action = SoftPinAction.ENABLE result = store.upsert_pub_tensor("my_tensor", tensor, config) if result == 0: diff --git a/docs/source/deployment/mooncake-store-deployment-guide.md b/docs/source/deployment/mooncake-store-deployment-guide.md index 640aec3c2e..d4d055ac58 100644 --- a/docs/source/deployment/mooncake-store-deployment-guide.md +++ b/docs/source/deployment/mooncake-store-deployment-guide.md @@ -647,6 +647,7 @@ mooncake_master \ |------|---------|-------------| | `--default_kv_lease_ttl` | `10000` ms | Lease TTL for KV objects. Supports `5000ms`, `5s`, `30m`, `1h` | | `--default_kv_soft_pin_ttl` | `1800000` ms | Soft pin TTL (30 min) | +| `--max_kv_soft_pin_ttl` | `86400000` ms | Maximum request-level soft pin TTL (24 h) | | `--allow_evict_soft_pinned_objects` | `true` | Allow evicting soft-pinned objects | | `--eviction_ratio` | `0.05` | Fraction evicted at high watermark | | `--eviction_high_watermark_ratio` | `0.90` | Usage ratio triggering eviction | diff --git a/docs/source/design/mooncake-store.md b/docs/source/design/mooncake-store.md index aa68f03cf8..dee3f9a651 100644 --- a/docs/source/design/mooncake-store.md +++ b/docs/source/design/mooncake-store.md @@ -478,7 +478,7 @@ On the Master side, group state is tenant-scoped. Objects with a non-empty group Group metadata affects lifecycle behavior on a best-effort basis: -- `ExistKey` and `GetReplicaList` refresh the lease, and the soft-pin timeout if present, for the current members of the group. +- `ExistKey` and `GetReplicaList` refresh the ordinary read lease for the current members of the group. Object soft-pin deadlines are independent and are not extended. - Memory eviction expands a grouped candidate to the group's current members and then applies the existing per-object safety checks. Members with active leases, hard pins, soft pins when soft-pin eviction is disabled, incomplete writes, busy replicas, or unavailable replica states are skipped. - Object removal APIs, copy/move tasks, and NoF eviction keep their existing object-level semantics. Group routing and membership metadata are cleaned up when objects are removed. @@ -659,14 +659,21 @@ The default lease TTL is 10 seconds and is configurable via a startup parameter For important and frequently used objects, such as system prompts, Mooncake Store provides a soft pin mechanism. When putting an object, it can be configured to enable soft pin. During eviction, objects that are not soft pinned are prioritized for eviction. Soft pinned objects are only evicted when memory is insufficient and no other objects are eligible for eviction. -If a soft pinned object is not accessed for an extended period, its soft pin status will be removed. If it is accessed again later, it will automatically be soft pinned once more. +The soft-pin lifetime starts when the first replica becomes readable. When its deadline is reached, the object becomes ordinary cache; later reads grant only an ordinary read lease and do not reactivate soft pinning. A later write can explicitly enable it again. -There are two startup parameters in `master_service` related to the soft pin mechanism: +Soft pin is runtime-only eviction-priority state. It is not persisted in snapshots or the HA OpLog, so recovery and Standby promotion downgrade restored objects to ordinary cache. Existing snapshot fields are retained only for format compatibility and ignored during recovery. -- `default_kv_soft_pin_ttl`: The duration (in milliseconds) after which a soft pinned object will have its soft pin status removed if not accessed. The default value is `30 minutes`. +There are three startup parameters in `master_service` related to the soft pin mechanism: + +- `default_kv_soft_pin_ttl`: The fixed soft-pin lifetime (in milliseconds) used when an `ENABLE` request omits `soft_pin_ttl_ms`. The default value is `30 minutes`; reads do not extend it. + +- `max_kv_soft_pin_ttl`: The largest request-level soft-pin TTL accepted by the Master. The default value is `24 hours`. - `allow_evict_soft_pinned_objects`: Whether soft pinned objects are allowed to be evicted. The default value is `true`. +An explicit `ENABLE` TTL of zero commits the object as ordinary cache. TTL +overrides are rejected for `PRESERVE` and `DISABLE`. + Notably, soft pinned objects can still be removed using APIs such as `Remove` or `RemoveAll`. ## Hard Pin @@ -677,9 +684,9 @@ Hard pin is set at object creation time through the `with_hard_pin` field in `Re Key differences from soft pin: -- Hard pin never expires. Soft pin status is removed after a configurable TTL if the object is not accessed. +- Hard pin never expires. Soft pin expires at a fixed deadline that starts when the first replica becomes readable, regardless of later accesses. - Hard-pinned objects are completely skipped during eviction. Soft-pinned objects may still be evicted when no other candidates are available. -- Hard pin is immutable once set. Soft pin status is automatically refreshed on access. +- Hard pin is immutable once set. Soft pin can be explicitly preserved, enabled, or disabled by write requests; reads do not refresh or reactivate it. ## Zombie Object Cleanup @@ -704,7 +711,8 @@ The preferred segment allocation feature is implemented through the `AllocationS ```cpp struct ReplicateConfig { size_t replica_num{1}; // Total number of replicas for the object - bool with_soft_pin{false}; // Whether to enable soft pin mechanism for this object + SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; + std::optional soft_pin_ttl_ms{}; // ENABLE override; omitted uses the Master default bool with_hard_pin{false}; // Whether to enable hard pin (never evicted) std::string preferred_segment{}; // Preferred segment for allocation }; diff --git a/docs/source/zh_archive/mooncake-store.md b/docs/source/zh_archive/mooncake-store.md index aade0b725d..b5cffe90d4 100644 --- a/docs/source/zh_archive/mooncake-store.md +++ b/docs/source/zh_archive/mooncake-store.md @@ -103,7 +103,8 @@ tl::expected Put(const ObjectKey& key, ```C++ struct ReplicateConfig { size_t replica_num{1}; // 对象的总副本数 - bool with_soft_pin{false}; // 是否为该对象启用软固定机制 + SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; // 软固定状态转换 + std::optional soft_pin_ttl_ms{}; // ENABLE 时可覆盖默认 TTL std::string preferred_segment{}; // 首选的分配段 }; ``` @@ -628,11 +629,15 @@ virtual tl::expected, ErrorCode> Allocate( 对于重要且频繁使用的对象,例如 system prompt,Mooncake Store 提供了软固定(soft pin)机制。在执行 `Put` 操作时,可以选择为特定的对象开启软固定机制。在执行替换任务时,系统会优先替换未被软固定的对象。仅当内存不足且没有其他对象可以被替换时,才会替换被软固定的对象。 -如果某个软固定的对象长时间未被访问,其软固定状态将被解除。而后当该对象再次被访问时,它将自动重新进入软固定状态。 +soft pin 生命周期从首个副本变为可读时开始。deadline 到达后,对象降级为普通 Cache;后续访问只授予普通读租约,不会重新启用 soft pin。后续写入仍可显式重新启用。 -`master_service` 中有两个与软固定机制相关的启动参数: +soft pin 是仅在运行时生效的淘汰优先级状态,不会持久化到快照或 HA OpLog。恢复或 Standby 提升后,恢复出的对象将降级为普通 Cache;快照中的兼容字段仅用于保持格式,恢复时会被忽略。 -* `default_kv_soft_pin_ttl`:表示一个被软固定的对象在多长时间(毫秒)未被访问后会自动解除软固定状态。默认值为`30 分钟`。 +`master_service` 中有三个与软固定机制相关的启动参数: + +* `default_kv_soft_pin_ttl`:未显式传入 TTL 时使用的固定软固定生命周期。默认值为 `30 分钟`,访问不会续期。 + +* `max_kv_soft_pin_ttl`:Master 接受的请求级 soft pin TTL 上限。默认值为 `24 小时`。 * `allow_evict_soft_pinned_objects`:是否允许替换已被软固定的对象。默认值为 `true`。 @@ -661,7 +666,8 @@ Mooncake Store 提供了**首选段分配**功能,允许用户为对象分配 ```cpp struct ReplicateConfig { size_t replica_num{1}; // 对象的总副本数 - bool with_soft_pin{false}; // 是否为该对象启用软固定机制 + SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; // 软固定状态转换 + std::optional soft_pin_ttl_ms{}; // ENABLE 时可覆盖默认 TTL std::string preferred_segment{}; // 首选的分配段 }; ``` diff --git a/mooncake-integration/store/store_py.cpp b/mooncake-integration/store/store_py.cpp index 1d239d70b1..fdeb35432d 100644 --- a/mooncake-integration/store/store_py.cpp +++ b/mooncake-integration/store/store_py.cpp @@ -1914,12 +1914,18 @@ PYBIND11_MODULE(store, m) { .value("GENERAL", ObjectDataType::GENERAL) .export_values(); + py::enum_(m, "SoftPinAction") + .value("PRESERVE", SoftPinAction::PRESERVE) + .value("ENABLE", SoftPinAction::ENABLE) + .value("DISABLE", SoftPinAction::DISABLE); + // Define the ReplicateConfig class py::class_(m, "ReplicateConfig") .def(py::init<>()) .def_readwrite("replica_num", &ReplicateConfig::replica_num) .def_readwrite("nof_replica_num", &ReplicateConfig::nof_replica_num) - .def_readwrite("with_soft_pin", &ReplicateConfig::with_soft_pin) + .def_readwrite("soft_pin_action", &ReplicateConfig::soft_pin_action) + .def_readwrite("soft_pin_ttl_ms", &ReplicateConfig::soft_pin_ttl_ms) .def_readwrite("with_hard_pin", &ReplicateConfig::with_hard_pin) .def_readwrite("preferred_segments", &ReplicateConfig::preferred_segments) diff --git a/mooncake-integration/store/store_py_internal.h b/mooncake-integration/store/store_py_internal.h index a7d030a4d7..e7f439f443 100644 --- a/mooncake-integration/store/store_py_internal.h +++ b/mooncake-integration/store/store_py_internal.h @@ -861,8 +861,10 @@ bool parallelism_specs_equal_by_kind(const TensorParallelismSpec &lhs, } bool is_default_replicate_config(const ReplicateConfig &config) { - return config.replica_num == 1 && !config.with_soft_pin && - !config.with_hard_pin && config.preferred_segments.empty() && + return config.replica_num == 1 && + config.soft_pin_action == SoftPinAction::PRESERVE && + !config.soft_pin_ttl_ms.has_value() && !config.with_hard_pin && + config.preferred_segments.empty() && config.preferred_segment.empty() && !config.prefer_alloc_in_same_node && !config.group_ids.has_value(); } diff --git a/mooncake-store/conf/master.json b/mooncake-store/conf/master.json index f55bdaf2b2..9d163ee627 100644 --- a/mooncake-store/conf/master.json +++ b/mooncake-store/conf/master.json @@ -9,6 +9,7 @@ "rpc_enable_tcp_no_delay": true, "default_kv_lease_ttl": 10000, "default_kv_soft_pin_ttl": 1800000, + "max_kv_soft_pin_ttl": 86400000, "allow_evict_soft_pinned_objects": true, "eviction_ratio": 0.1, "eviction_high_watermark_ratio": 1.0, diff --git a/mooncake-store/conf/master.yaml b/mooncake-store/conf/master.yaml index 545807008f..c970f61676 100644 --- a/mooncake-store/conf/master.yaml +++ b/mooncake-store/conf/master.yaml @@ -9,6 +9,7 @@ rpc_enable_tcp_no_delay: true default_kv_lease_ttl: 10000 default_kv_soft_pin_ttl: 1800000 +max_kv_soft_pin_ttl: 86400000 allow_evict_soft_pinned_objects: true eviction_ratio: 0.1 # Overrides the 0.90 code default. A value of 1.0 disables proactive diff --git a/mooncake-store/include/master_config.h b/mooncake-store/include/master_config.h index 4aa71b3fac..eae4cec3a0 100644 --- a/mooncake-store/include/master_config.h +++ b/mooncake-store/include/master_config.h @@ -40,6 +40,7 @@ struct MasterConfig { uint64_t default_kv_lease_ttl; uint64_t default_kv_soft_pin_ttl; + uint64_t max_kv_soft_pin_ttl = DEFAULT_MAX_KV_SOFT_PIN_TTL_MS; bool allow_evict_soft_pinned_objects; double eviction_ratio; double eviction_high_watermark_ratio; @@ -182,6 +183,7 @@ class MasterServiceSupervisorConfig { RequiredParam rpc_thread_num{"rpc_thread_num"}; // Parameters with default values (optional parameters) + uint64_t max_kv_soft_pin_ttl = DEFAULT_MAX_KV_SOFT_PIN_TTL_MS; std::string rpc_address = "0.0.0.0"; std::string metrics_host = "0.0.0.0"; std::chrono::steady_clock::duration rpc_conn_timeout = std::chrono::seconds( @@ -274,6 +276,7 @@ class MasterServiceSupervisorConfig { metrics_host = config.metrics_host; default_kv_lease_ttl = config.default_kv_lease_ttl; default_kv_soft_pin_ttl = config.default_kv_soft_pin_ttl; + max_kv_soft_pin_ttl = config.max_kv_soft_pin_ttl; allow_evict_soft_pinned_objects = config.allow_evict_soft_pinned_objects; eviction_ratio = config.eviction_ratio; @@ -456,6 +459,7 @@ class WrappedMasterServiceConfig { // Optional parameters (with default values) uint64_t default_kv_soft_pin_ttl = DEFAULT_KV_SOFT_PIN_TTL_MS; + uint64_t max_kv_soft_pin_ttl = DEFAULT_MAX_KV_SOFT_PIN_TTL_MS; bool allow_evict_soft_pinned_objects = DEFAULT_ALLOW_EVICT_SOFT_PINNED_OBJECTS; bool enable_metric_reporting = true; @@ -547,6 +551,7 @@ class WrappedMasterServiceConfig { // Set optional parameters (these have default values) default_kv_soft_pin_ttl = config.default_kv_soft_pin_ttl; + max_kv_soft_pin_ttl = config.max_kv_soft_pin_ttl; allow_evict_soft_pinned_objects = config.allow_evict_soft_pinned_objects; enable_metric_reporting = config.enable_metric_reporting; @@ -662,6 +667,7 @@ class WrappedMasterServiceConfig { // Set optional parameters (these have default values) default_kv_soft_pin_ttl = config.default_kv_soft_pin_ttl; + max_kv_soft_pin_ttl = config.max_kv_soft_pin_ttl; allow_evict_soft_pinned_objects = config.allow_evict_soft_pinned_objects; enable_metric_reporting = config.enable_metric_reporting; @@ -751,6 +757,7 @@ class MasterServiceConfigBuilder { private: uint64_t default_kv_lease_ttl_ = DEFAULT_DEFAULT_KV_LEASE_TTL; uint64_t default_kv_soft_pin_ttl_ = DEFAULT_KV_SOFT_PIN_TTL_MS; + uint64_t max_kv_soft_pin_ttl_ = DEFAULT_MAX_KV_SOFT_PIN_TTL_MS; bool allow_evict_soft_pinned_objects_ = DEFAULT_ALLOW_EVICT_SOFT_PINNED_OBJECTS; double eviction_ratio_ = DEFAULT_EVICTION_RATIO; @@ -821,6 +828,11 @@ class MasterServiceConfigBuilder { return *this; } + MasterServiceConfigBuilder& set_max_kv_soft_pin_ttl(uint64_t ttl) { + max_kv_soft_pin_ttl_ = ttl; + return *this; + } + MasterServiceConfigBuilder& set_allow_evict_soft_pinned_objects( bool allow) { allow_evict_soft_pinned_objects_ = allow; @@ -1108,6 +1120,7 @@ class MasterServiceConfig { public: uint64_t default_kv_lease_ttl = DEFAULT_DEFAULT_KV_LEASE_TTL; uint64_t default_kv_soft_pin_ttl = DEFAULT_KV_SOFT_PIN_TTL_MS; + uint64_t max_kv_soft_pin_ttl = DEFAULT_MAX_KV_SOFT_PIN_TTL_MS; bool allow_evict_soft_pinned_objects = DEFAULT_ALLOW_EVICT_SOFT_PINNED_OBJECTS; double eviction_ratio = DEFAULT_EVICTION_RATIO; @@ -1195,6 +1208,7 @@ class MasterServiceConfig { default_kv_lease_ttl = config.default_kv_lease_ttl; default_kv_soft_pin_ttl = config.default_kv_soft_pin_ttl; + max_kv_soft_pin_ttl = config.max_kv_soft_pin_ttl; allow_evict_soft_pinned_objects = config.allow_evict_soft_pinned_objects; eviction_ratio = config.eviction_ratio; @@ -1285,6 +1299,7 @@ inline MasterServiceConfig MasterServiceConfigBuilder::build() const { MasterServiceConfig config; config.default_kv_lease_ttl = default_kv_lease_ttl_; config.default_kv_soft_pin_ttl = default_kv_soft_pin_ttl_; + config.max_kv_soft_pin_ttl = max_kv_soft_pin_ttl_; config.allow_evict_soft_pinned_objects = allow_evict_soft_pinned_objects_; config.eviction_ratio = eviction_ratio_; config.eviction_high_watermark_ratio = eviction_high_watermark_ratio_; diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index 93cea6fe6a..d341993059 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include #include @@ -9,9 +10,11 @@ #include #include #include +#include #include #include #include +#include #include #include #include @@ -68,6 +71,7 @@ struct MetadataStoragePlugin; // Forward declarations for test classes namespace test { +class MasterServiceTest; class MasterServiceSnapshotTestBase; class SnapshotChildProcessTest; // Friended so the promotion-on-hit tests can drive a serialize/reset/ @@ -100,6 +104,7 @@ class BatchEvictBench; * 4. metadata_shards_[shard_idx_].mutex * 5. tenant_quota_recompute_mutex_ * 6. ShardedTenantQuotaTable internal mutex or segment_mutex_ + * 7. soft_pin_deadline_index_ mutex * * Strict tenant admission and policy mutation paths that need both * tenant_quota_policy_mutex_ and snapshot_mutex_ must acquire the tenant @@ -111,6 +116,7 @@ class BatchEvictBench; class MasterService { // Test friend class for snapshot/restore testing friend class test::MasterServiceSnapshotTestBase; + friend class test::MasterServiceTest; friend class test::SnapshotChildProcessTest; friend class test::PromotionOnHitTest; friend class benchmarks::BatchEvictBench; @@ -933,7 +939,27 @@ class MasterService { std::string user_key; }; + struct ResolvedSoftPinRequest { + SoftPinAction action{SoftPinAction::PRESERVE}; + uint64_t ttl_ms{0}; + }; + struct ObjectMetadata { + struct SoftPinEvaluation { + bool active{false}; + int metric_delta{0}; + std::optional + removed_deadline; + std::optional + deadline_to_index; + }; + + struct PendingSoftPinAction { + SoftPinAction action{SoftPinAction::PRESERVE}; + uint64_t ttl_ms{0}; + std::vector eligible_replica_ids; + }; + // RAII-style metric management ~ObjectMetadata() { MasterMetricManager::instance().dec_key_count(1); @@ -948,7 +974,9 @@ class MasterService { const UUID& client_id_, const std::chrono::system_clock::time_point put_start_time_, size_t value_length, std::vector&& reps, - bool enable_soft_pin, bool enable_hard_pin = false, + std::optional + committed_soft_pin_timeout = std::nullopt, + bool enable_hard_pin = false, ObjectDataType data_type_ = ObjectDataType::UNKNOWN, std::string group_id_ = "", TenantId tenant_id_ = TenantId(), std::string user_key_ = {}) @@ -960,12 +988,11 @@ class MasterService { tenant_id(std::move(tenant_id_)), user_key(std::move(user_key_)), lease_timeout(), - soft_pin_timeout(std::nullopt), + soft_pin_timeout(std::move(committed_soft_pin_timeout)), hard_pinned(enable_hard_pin), replicas_(std::move(reps)) { MasterMetricManager::instance().inc_key_count(1); - if (enable_soft_pin) { - soft_pin_timeout.emplace(); + if (soft_pin_timeout) { MasterMetricManager::instance().inc_soft_pin_key_count(1); } MasterMetricManager::instance().observe_value_size(value_length); @@ -993,9 +1020,13 @@ class MasterService { mutable std::chrono::system_clock::time_point lease_timeout GUARDED_BY(lock); // hard lease mutable std::optional - soft_pin_timeout GUARDED_BY(lock); // optional soft pin, only - // set for vip objects - const bool hard_pinned{false}; // immutable, set at creation + soft_pin_timeout GUARDED_BY(lock); // committed object soft-pin + // deadline + // Replica IDs scope this action to the current write. PutEnd does not + // carry a generation token, so a stale End from the same client cannot + // otherwise be distinguished from the current write. + std::optional pending_soft_pin_action; + const bool hard_pinned{false}; // immutable, set at creation bool memory_cache_total_accounted{false}; bool disk_cache_total_accounted{false}; uint64_t reserved_quota_charge_bytes{0}; @@ -1155,31 +1186,20 @@ class MasterService { }); } - // Grant a lease with timeout as now() + ttl, only update if the new - // timeout is larger - void GrantLease(const uint64_t ttl, const uint64_t soft_ttl) const { + // Grant an ordinary read lease with timeout as now() + ttl. Soft-pin + // lifetime is intentionally independent from reads. + void GrantReadLease(const uint64_t ttl) const { SpinLocker locker(&lock); std::chrono::system_clock::time_point now = std::chrono::system_clock::now(); lease_timeout = std::max(lease_timeout, now + std::chrono::milliseconds(ttl)); - if (soft_pin_timeout) { - soft_pin_timeout = - std::max(*soft_pin_timeout, - now + std::chrono::milliseconds(soft_ttl)); - } } - bool NeedsLeaseRefresh(const uint64_t ttl, - const uint64_t soft_ttl) const { + bool NeedsReadLeaseRefresh(const uint64_t ttl) const { SpinLocker locker(&lock); const auto now = std::chrono::system_clock::now(); - if (lease_timeout <= now + std::chrono::milliseconds(ttl / 2)) { - return true; - } - return soft_pin_timeout && - *soft_pin_timeout <= - now + std::chrono::milliseconds(soft_ttl / 2); + return lease_timeout <= now + std::chrono::milliseconds(ttl / 2); } // Check if the lease has expired @@ -1194,17 +1214,159 @@ class MasterService { return now >= lease_timeout; } - // Check if is in soft pin status - bool IsSoftPinned() const { + SoftPinEvaluation EvaluateSoftPin( + const std::chrono::system_clock::time_point& now) const { + SpinLocker locker(&lock); + if (soft_pin_timeout && now >= *soft_pin_timeout) { + const auto removed_deadline = *soft_pin_timeout; + soft_pin_timeout.reset(); + return {.active = false, + .metric_delta = -1, + .removed_deadline = removed_deadline, + .deadline_to_index = std::nullopt}; + } + return {.active = soft_pin_timeout.has_value(), + .metric_delta = 0, + .removed_deadline = std::nullopt, + .deadline_to_index = std::nullopt}; + } + + bool ExpireSoftPinIfDeadlineMatches( + const std::chrono::system_clock::time_point& expected_deadline, + const std::chrono::system_clock::time_point& now) const { + SpinLocker locker(&lock); + if (!soft_pin_timeout || *soft_pin_timeout != expected_deadline || + now < expected_deadline) { + return false; + } + soft_pin_timeout.reset(); + return true; + } + + static std::chrono::system_clock::time_point ComputeSoftPinDeadline( + const std::chrono::system_clock::time_point& now, uint64_t ttl_ms) { + using Milliseconds = std::chrono::milliseconds; + using MillisecondsRep = Milliseconds::rep; + const auto max_time = std::chrono::system_clock::time_point::max(); + if (ttl_ms > static_cast( + std::numeric_limits::max())) { + return max_time; + } + const auto remaining_ms = + std::chrono::duration_cast(max_time - now) + .count(); + if (remaining_ms < 0 || + ttl_ms > static_cast(remaining_ms)) { + return max_time; + } + const auto ttl = Milliseconds(static_cast(ttl_ms)); + return now + ttl; + } + + std::optional + GetCommittedSoftPinTimeout() const { SpinLocker locker(&lock); - return soft_pin_timeout && - std::chrono::system_clock::now() < *soft_pin_timeout; + return soft_pin_timeout; + } + + void BeginSoftPinAction(const ResolvedSoftPinRequest& request, + std::vector eligible_replica_ids) { + pending_soft_pin_action = + PendingSoftPinAction{request.action, request.ttl_ms, + std::move(eligible_replica_ids)}; } - // Check if is in soft pin status - bool IsSoftPinned(std::chrono::system_clock::time_point& now) const { + bool PendingSoftPinOwnsReplica(ReplicaID replica_id) const { + if (!pending_soft_pin_action) { + return false; + } + const auto& eligible = + pending_soft_pin_action->eligible_replica_ids; + return std::find(eligible.begin(), eligible.end(), replica_id) != + eligible.end(); + } + + void ClearPendingSoftPinAction() { pending_soft_pin_action.reset(); } + + SoftPinEvaluation CommitPendingSoftPin( + const std::chrono::system_clock::time_point& now) { + if (!pending_soft_pin_action) { + return EvaluateSoftPin(now); + } + + const PendingSoftPinAction pending = + std::move(*pending_soft_pin_action); + pending_soft_pin_action.reset(); + SpinLocker locker(&lock); - return soft_pin_timeout && now < *soft_pin_timeout; + int metric_delta = 0; + std::optional + removed_deadline; + std::optional + deadline_to_index; + if (soft_pin_timeout && now >= *soft_pin_timeout) { + removed_deadline = *soft_pin_timeout; + soft_pin_timeout.reset(); + --metric_delta; + } + + switch (pending.action) { + case SoftPinAction::PRESERVE: + break; + case SoftPinAction::ENABLE: + if (pending.ttl_ms == 0) { + if (soft_pin_timeout) { + removed_deadline = *soft_pin_timeout; + soft_pin_timeout.reset(); + --metric_delta; + } + } else { + if (!soft_pin_timeout) { + ++metric_delta; + } + soft_pin_timeout = + ComputeSoftPinDeadline(now, pending.ttl_ms); + deadline_to_index = soft_pin_timeout; + // Upserting the latest registration supersedes any + // expired or previously active deadline. + removed_deadline.reset(); + } + break; + case SoftPinAction::DISABLE: + if (soft_pin_timeout) { + removed_deadline = *soft_pin_timeout; + soft_pin_timeout.reset(); + --metric_delta; + } + break; + } + return {.active = soft_pin_timeout.has_value(), + .metric_delta = metric_delta, + .removed_deadline = removed_deadline, + .deadline_to_index = deadline_to_index}; + } + + void ClearPendingSoftPinIfNoViableReplica() { + if (!pending_soft_pin_action) { + return; + } + const auto& eligible = + pending_soft_pin_action->eligible_replica_ids; + const bool has_viable_replica = + std::any_of(replicas_.begin(), replicas_.end(), + [&eligible](const Replica& replica) { + const bool belongs_to_write = + std::find(eligible.begin(), eligible.end(), + replica.id()) != eligible.end(); + const bool valid_handle = + !replica.has_invalid_mem_handle() && + !replica.has_invalid_nof_handle(); + return belongs_to_write && + replica.is_processing() && valid_handle; + }); + if (!has_viable_replica) { + pending_soft_pin_action.reset(); + } } bool IsHardPinned() const { return hard_pinned; } @@ -1352,6 +1514,53 @@ class MasterService { }; std::array metadata_shards_; + class SoftPinDeadlineIndex { + public: + using TimePoint = std::chrono::system_clock::time_point; + + struct Entry { + TimePoint deadline; + size_t shard_idx; + std::string scoped_key; + }; + + void Upsert(std::string scoped_key, size_t shard_idx, + const TimePoint& deadline); + void Remove(const std::string& scoped_key); + void RemoveIfMatches(const std::string& scoped_key, size_t shard_idx, + const TimePoint& deadline); + std::vector PopExpired(const TimePoint& now); + void Clear(); + + size_t HeapSizeForTest() const; + size_t RegistrationCountForTest() const; + + private: + struct Registration { + TimePoint deadline; + size_t shard_idx; + }; + + struct EarlierDeadline { + bool operator()(const Entry& lhs, const Entry& rhs) const { + return lhs.deadline > rhs.deadline; + } + }; + + static constexpr size_t kMinCompactionThreshold = 4096; + static constexpr size_t kCompactionRatio = 2; + + void MaybeCompactLocked() REQUIRES(mutex_); + + mutable std::mutex mutex_; + std::priority_queue, EarlierDeadline> heap_ + GUARDED_BY(mutex_); + std::unordered_map registrations_ + GUARDED_BY(mutex_); + }; + + mutable SoftPinDeadlineIndex soft_pin_deadline_index_; + static bool HasCompletedMemoryCacheReplica(const ObjectMetadata& metadata); static bool HasCompletedDiskCacheReplica(const ObjectMetadata& metadata); static void SyncCacheTotalAccounting(ObjectMetadata& metadata); @@ -1558,6 +1767,18 @@ class MasterService { void GrantLeaseForGroup(const TenantState& tenant_state, const std::string& key, const ObjectMetadata& metadata) const; + static void ApplySoftPinMetricDelta(int metric_delta); + size_t GetMetadataShardIndex(const ObjectMetadata& metadata) const; + void ApplySoftPinEvaluation( + const ObjectMetadata& metadata, + const ObjectMetadata::SoftPinEvaluation& result) const; + bool IsSoftPinActive( + const ObjectMetadata& metadata, + const std::chrono::system_clock::time_point& now) const; + void CleanupExpiredSoftPins( + const std::chrono::system_clock::time_point& now); + auto ResolveSoftPinRequest(const ReplicateConfig& config) const + -> tl::expected; // Helper to clean up stale handles pointing to unmounted segments // or local_disk replicas whose owner client is no longer alive. @@ -1573,7 +1794,10 @@ class MasterService { const std::string& key, uint64_t value_length, const ReplicateConfig& config, const std::string& group_id, const TenantId& tenant_id, - const std::chrono::system_clock::time_point& now) + const std::chrono::system_clock::time_point& now, + const ResolvedSoftPinRequest& soft_pin_request, + std::optional + committed_soft_pin_timeout = std::nullopt) -> tl::expected, ErrorCode>; /** @@ -1671,6 +1895,7 @@ class MasterService { // Lease related members const uint64_t default_kv_lease_ttl_; // in milliseconds const uint64_t default_kv_soft_pin_ttl_; // in milliseconds + const uint64_t max_kv_soft_pin_ttl_; // in milliseconds const bool allow_evict_soft_pinned_objects_; // Eviction related members @@ -1827,8 +2052,7 @@ class MasterService { } void Create(const UUID& client_id, uint64_t total_length, - std::vector replicas, bool enable_soft_pin, - bool enable_hard_pin = false, + std::vector replicas, bool enable_hard_pin = false, ObjectDataType data_type = ObjectDataType::UNKNOWN, std::string group_id = "") { if (Exists()) { @@ -1841,7 +2065,7 @@ class MasterService { std::forward_as_tuple(object_id_.user_key), std::forward_as_tuple( client_id, now, total_length, std::move(replicas), - enable_soft_pin, enable_hard_pin, data_type, group_id, + std::nullopt, enable_hard_pin, data_type, group_id, object_id_.tenant_id, object_id_.user_key)); it_ = result.first; if (result.second) { diff --git a/mooncake-store/include/metadata_store.h b/mooncake-store/include/metadata_store.h index 97c886ddf7..8eebcf52a9 100644 --- a/mooncake-store/include/metadata_store.h +++ b/mooncake-store/include/metadata_store.h @@ -31,6 +31,8 @@ struct StandbyObjectMetadata { // 1. Standby does not perform eviction, so lease info is not used // 2. After promotion, new Primary should grant fresh leases, not restore // old ones + // Soft pin is also omitted because it is runtime-only eviction-priority + // state; promoted objects resume as ordinary cache. uint64_t last_sequence_id{ 0}; // Last OpLog sequence ID that modified this key std::string group_id; // Tenant group identifier diff --git a/mooncake-store/include/replica.h b/mooncake-store/include/replica.h index 3aa42ecf98..49696e5e50 100644 --- a/mooncake-store/include/replica.h +++ b/mooncake-store/include/replica.h @@ -57,6 +57,28 @@ enum class ReplicaStatus { FAILED, // Failed state (can be used for reassignment) }; +/** + * @brief Requested soft-pin transition for a Put or Upsert operation. + */ +enum class SoftPinAction : uint8_t { + PRESERVE = 0, + ENABLE = 1, + DISABLE = 2, +}; + +inline std::ostream& operator<<(std::ostream& os, + const SoftPinAction& action) noexcept { + switch (action) { + case SoftPinAction::PRESERVE: + return os << "PRESERVE"; + case SoftPinAction::ENABLE: + return os << "ENABLE"; + case SoftPinAction::DISABLE: + return os << "DISABLE"; + } + return os << "UNKNOWN"; +} + /** * @brief Stream operator for ReplicaStatus */ @@ -81,7 +103,10 @@ inline std::ostream& operator<<(std::ostream& os, struct ReplicateConfig { size_t replica_num{1}; size_t nof_replica_num{0}; - bool with_soft_pin{false}; + SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; + // Optional request-level override. When omitted, ENABLE uses the + // master's default soft-pin TTL. + std::optional soft_pin_ttl_ms{}; bool with_hard_pin{false}; // Hard pin: object cannot be evicted std::vector preferred_segments{}; // Preferred segments for allocation @@ -110,8 +135,14 @@ struct ReplicateConfig { const ReplicateConfig& config) noexcept { os << "ReplicateConfig: { replica_num: " << config.replica_num << ", nof_replica_num: " << config.nof_replica_num - << ", with_soft_pin: " << config.with_soft_pin - << ", with_hard_pin: " << config.with_hard_pin + << ", soft_pin_action: " << config.soft_pin_action + << ", soft_pin_ttl_ms: "; + if (config.soft_pin_ttl_ms.has_value()) { + os << *config.soft_pin_ttl_ms; + } else { + os << "default"; + } + os << ", with_hard_pin: " << config.with_hard_pin << ", preferred_segments: ["; for (size_t i = 0; i < config.preferred_segments.size(); ++i) { os << config.preferred_segments[i]; diff --git a/mooncake-store/include/types.h b/mooncake-store/include/types.h index c85eb8288c..0721098bbb 100644 --- a/mooncake-store/include/types.h +++ b/mooncake-store/include/types.h @@ -88,6 +88,8 @@ static constexpr uint64_t DEFAULT_DEFAULT_KV_LEASE_TTL = 10000; // in milliseconds static constexpr uint64_t DEFAULT_KV_SOFT_PIN_TTL_MS = 30 * 60 * 1000; // 30 minutes +static constexpr uint64_t DEFAULT_MAX_KV_SOFT_PIN_TTL_MS = + 24 * 60 * 60 * 1000; // 24 hours static constexpr bool DEFAULT_ALLOW_EVICT_SOFT_PINNED_OBJECTS = true; static constexpr double DEFAULT_EVICTION_RATIO = 0.05; static constexpr double DEFAULT_EVICTION_HIGH_WATERMARK_RATIO = 0.90; diff --git a/mooncake-store/src/ha/snapshot/catalog_backed_snapshot_provider.cpp b/mooncake-store/src/ha/snapshot/catalog_backed_snapshot_provider.cpp index 54ee82056d..d2eec6c408 100644 --- a/mooncake-store/src/ha/snapshot/catalog_backed_snapshot_provider.cpp +++ b/mooncake-store/src/ha/snapshot/catalog_backed_snapshot_provider.cpp @@ -127,10 +127,21 @@ DeserializeStandbyObjectMetadata( (void)array[index++].as(); // put_start_time const auto size = static_cast(array[index++].as()); const auto lease_timestamp_ms = array[index++].as(); - const bool has_soft_pin_timeout = array[index++].as(); - const auto soft_pin_timestamp_ms = array[index++].as(); + (void)array[index++].as(); // legacy soft-pin flag + (void)array[index++].as(); // legacy soft-pin deadline const auto replica_count = array[index++].as(); + const auto max_timestamp_ms = + std::chrono::duration_cast( + std::chrono::system_clock::time_point::max().time_since_epoch()) + .count(); + if (max_timestamp_ms < 0 || + lease_timestamp_ms > static_cast(max_timestamp_ms)) { + LOG(ERROR) << "Snapshot metadata timestamp exceeds system_clock " + "range"; + return tl::make_unexpected(ErrorCode::DESERIALIZE_FAIL); + } + // Optional fields are decoded by type for backward/forward // compatibility with MasterService::MetadataSerializer, which appends // them over time: @@ -167,15 +178,7 @@ DeserializeStandbyObjectMetadata( const auto lease_timeout = std::chrono::system_clock::time_point( std::chrono::milliseconds(lease_timestamp_ms)); - std::optional soft_pin_timeout; - if (has_soft_pin_timeout) { - soft_pin_timeout.emplace( - std::chrono::milliseconds(soft_pin_timestamp_ms)); - } - - if (size == 0 || - (lease_timeout <= now && (!soft_pin_timeout.has_value() || - soft_pin_timeout.value() <= now))) { + if (size == 0 || lease_timeout <= now) { return std::optional(); } diff --git a/mooncake-store/src/master.cpp b/mooncake-store/src/master.cpp index c9fa9b5a94..42d954c115 100644 --- a/mooncake-store/src/master.cpp +++ b/mooncake-store/src/master.cpp @@ -36,9 +36,13 @@ static_assert(mooncake::DEFAULT_DEFAULT_KV_LEASE_TTL == 10000, static_assert(mooncake::DEFAULT_KV_SOFT_PIN_TTL_MS == 30 * 60 * 1000, "Update kDefaultKvSoftPinTtlFlagValue when " "DEFAULT_KV_SOFT_PIN_TTL_MS changes"); +static_assert(mooncake::DEFAULT_MAX_KV_SOFT_PIN_TTL_MS == 24 * 60 * 60 * 1000, + "Update kDefaultMaxKvSoftPinTtlFlagValue when " + "DEFAULT_MAX_KV_SOFT_PIN_TTL_MS changes"); constexpr char kDefaultKvLeaseTtlFlagValue[] = "10000"; constexpr char kDefaultKvSoftPinTtlFlagValue[] = "1800000"; +constexpr char kDefaultMaxKvSoftPinTtlFlagValue[] = "86400000"; namespace { @@ -127,11 +131,16 @@ DEFINE_string(default_kv_lease_ttl, kDefaultKvLeaseTtlFlagValue, DEFINE_string(default_kv_soft_pin_ttl, kDefaultKvSoftPinTtlFlagValue, "Default soft pin TTL for kv objects. Supports raw milliseconds " "or duration strings with ms, s, m, or h suffixes"); +DEFINE_string(max_kv_soft_pin_ttl, kDefaultMaxKvSoftPinTtlFlagValue, + "Maximum request-level soft pin TTL for kv objects. Supports " + "raw milliseconds or duration strings with ms, s, m, or h " + "suffixes"); DEFINE_bool(allow_evict_soft_pinned_objects, mooncake::DEFAULT_ALLOW_EVICT_SOFT_PINNED_OBJECTS, "Whether to allow eviction of soft pinned objects during eviction"); DEFINE_validator(default_kv_lease_ttl, ValidateDurationFlag); DEFINE_validator(default_kv_soft_pin_ttl, ValidateDurationFlag); +DEFINE_validator(max_kv_soft_pin_ttl, ValidateDurationFlag); DEFINE_double(eviction_ratio, mooncake::DEFAULT_EVICTION_RATIO, "Ratio of objects to evict when Memory space is full"); DEFINE_double(eviction_high_watermark_ratio, @@ -465,6 +474,9 @@ void InitMasterConf(const mooncake::DefaultConfig& default_config, default_config.GetDurationMs("default_kv_soft_pin_ttl", &master_config.default_kv_soft_pin_ttl, mooncake::DEFAULT_KV_SOFT_PIN_TTL_MS); + default_config.GetDurationMs("max_kv_soft_pin_ttl", + &master_config.max_kv_soft_pin_ttl, + mooncake::DEFAULT_MAX_KV_SOFT_PIN_TTL_MS); default_config.GetBool("allow_evict_soft_pinned_objects", &master_config.allow_evict_soft_pinned_objects, FLAGS_allow_evict_soft_pinned_objects); @@ -802,6 +814,12 @@ void LoadConfigFromCmdline(mooncake::MasterConfig& master_config, master_config.default_kv_soft_pin_ttl = ParseDurationFlagOrDie( "default_kv_soft_pin_ttl", FLAGS_default_kv_soft_pin_ttl); } + if ((google::GetCommandLineFlagInfo("max_kv_soft_pin_ttl", &info) && + !info.is_default) || + !conf_set) { + master_config.max_kv_soft_pin_ttl = ParseDurationFlagOrDie( + "max_kv_soft_pin_ttl", FLAGS_max_kv_soft_pin_ttl); + } if ((google::GetCommandLineFlagInfo("allow_evict_soft_pinned_objects", &info) && !info.is_default) || @@ -1434,6 +1452,7 @@ int main(int argc, char* argv[]) { << ", metrics_host=" << master_config.metrics_host << ", default_kv_lease_ttl=" << master_config.default_kv_lease_ttl << ", default_kv_soft_pin_ttl=" << master_config.default_kv_soft_pin_ttl + << ", max_kv_soft_pin_ttl=" << master_config.max_kv_soft_pin_ttl << ", allow_evict_soft_pinned_objects=" << master_config.allow_evict_soft_pinned_objects << ", eviction_ratio=" << master_config.eviction_ratio diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index e0e30356c1..05494f793c 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -165,6 +165,7 @@ MasterService::MasterService(const MasterServiceConfig& config) }), default_kv_lease_ttl_(config.default_kv_lease_ttl), default_kv_soft_pin_ttl_(config.default_kv_soft_pin_ttl), + max_kv_soft_pin_ttl_(config.max_kv_soft_pin_ttl), allow_evict_soft_pinned_objects_(config.allow_evict_soft_pinned_objects), eviction_ratio_(config.eviction_ratio), eviction_high_watermark_ratio_(config.eviction_high_watermark_ratio), @@ -201,6 +202,13 @@ MasterService::MasterService(const MasterServiceConfig& config) offloading_queue_limit_(config.offloading_queue_limit), offload_cap_ratio_(config.offload_cap_ratio), task_manager_(config.task_manager_config) { + if (default_kv_soft_pin_ttl_ > max_kv_soft_pin_ttl_) { + LOG(ERROR) << "Invalid soft-pin TTL configuration: default=" + << default_kv_soft_pin_ttl_ + << ", max=" << max_kv_soft_pin_ttl_; + throw std::invalid_argument("Invalid soft-pin TTL configuration"); + } + // Initialize HTTP metadata key prefix (read env var once at startup) const char* custom_prefix = std::getenv("MC_METADATA_CLUSTER_ID"); if (custom_prefix && std::strlen(custom_prefix) > 0) { @@ -1921,6 +1929,9 @@ MasterService::EraseMetadata( ReleaseLocalDiskUsage(metadata.GetAllReplicas()); AccountCacheTotalRemoval(metadata); + if (metadata.GetCommittedSoftPinTimeout()) { + soft_pin_deadline_index_.Remove(tenant_id.MakeScopedKey(key)); + } switch (quota_mode) { case QuotaEraseMode::kFull: AbortTenantQuota(tenant_id, metadata.reserved_quota_charge_bytes); @@ -2000,16 +2011,224 @@ void MasterService::RebuildGroupRoutingIndex() { } } +void MasterService::SoftPinDeadlineIndex::MaybeCompactLocked() { + const size_t live_count = registrations_.size(); + const size_t ratio_limit = + live_count > std::numeric_limits::max() / kCompactionRatio + ? std::numeric_limits::max() + : live_count * kCompactionRatio; + const size_t threshold = std::max(kMinCompactionThreshold, ratio_limit); + if (heap_.size() <= threshold) { + return; + } + + decltype(heap_) rebuilt; + for (const auto& [scoped_key, registration] : registrations_) { + rebuilt.push( + Entry{registration.deadline, registration.shard_idx, scoped_key}); + } + heap_.swap(rebuilt); +} + +void MasterService::SoftPinDeadlineIndex::Upsert(std::string scoped_key, + size_t shard_idx, + const TimePoint& deadline) { + std::lock_guard lock(mutex_); + const auto it = registrations_.find(scoped_key); + if (it != registrations_.end() && it->second.deadline == deadline && + it->second.shard_idx == shard_idx) { + return; + } + + auto [registration_it, inserted] = registrations_.insert_or_assign( + scoped_key, Registration{deadline, shard_idx}); + (void)inserted; + heap_.push(Entry{deadline, shard_idx, registration_it->first}); + MaybeCompactLocked(); +} + +void MasterService::SoftPinDeadlineIndex::Remove( + const std::string& scoped_key) { + std::lock_guard lock(mutex_); + if (registrations_.erase(scoped_key) > 0) { + MaybeCompactLocked(); + } +} + +void MasterService::SoftPinDeadlineIndex::RemoveIfMatches( + const std::string& scoped_key, size_t shard_idx, + const TimePoint& deadline) { + std::lock_guard lock(mutex_); + const auto it = registrations_.find(scoped_key); + if (it != registrations_.end() && it->second.deadline == deadline && + it->second.shard_idx == shard_idx) { + registrations_.erase(it); + MaybeCompactLocked(); + } +} + +std::vector +MasterService::SoftPinDeadlineIndex::PopExpired(const TimePoint& now) { + std::vector expired; + std::lock_guard lock(mutex_); + while (!heap_.empty() && heap_.top().deadline <= now) { + Entry entry = heap_.top(); + heap_.pop(); + + const auto it = registrations_.find(entry.scoped_key); + if (it == registrations_.end() || + it->second.deadline != entry.deadline || + it->second.shard_idx != entry.shard_idx) { + continue; + } + registrations_.erase(it); + expired.push_back(std::move(entry)); + } + MaybeCompactLocked(); + return expired; +} + +void MasterService::SoftPinDeadlineIndex::Clear() { + std::lock_guard lock(mutex_); + registrations_.clear(); + decltype(heap_) empty; + heap_.swap(empty); +} + +size_t MasterService::SoftPinDeadlineIndex::HeapSizeForTest() const { + std::lock_guard lock(mutex_); + return heap_.size(); +} + +size_t MasterService::SoftPinDeadlineIndex::RegistrationCountForTest() const { + std::lock_guard lock(mutex_); + return registrations_.size(); +} + +auto MasterService::ResolveSoftPinRequest(const ReplicateConfig& config) const + -> tl::expected { + switch (config.soft_pin_action) { + case SoftPinAction::PRESERVE: + case SoftPinAction::DISABLE: + if (config.soft_pin_ttl_ms.has_value()) { + LOG(ERROR) << "soft_pin_action=" << config.soft_pin_action + << ", soft_pin_ttl_ms=" << *config.soft_pin_ttl_ms + << ", error=ttl_requires_enable"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + return ResolvedSoftPinRequest{config.soft_pin_action, 0}; + case SoftPinAction::ENABLE: { + const uint64_t ttl_ms = + config.soft_pin_ttl_ms.value_or(default_kv_soft_pin_ttl_); + if (ttl_ms > max_kv_soft_pin_ttl_) { + LOG(ERROR) << "soft_pin_ttl_ms=" << ttl_ms + << ", max_kv_soft_pin_ttl=" << max_kv_soft_pin_ttl_ + << ", error=soft_pin_ttl_exceeds_limit"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + return ResolvedSoftPinRequest{config.soft_pin_action, ttl_ms}; + } + } + LOG(ERROR) << "soft_pin_action=" + << static_cast(config.soft_pin_action) + << ", error=invalid_soft_pin_action"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); +} + +void MasterService::ApplySoftPinMetricDelta(int metric_delta) { + if (metric_delta > 0) { + MasterMetricManager::instance().inc_soft_pin_key_count(metric_delta); + } else if (metric_delta < 0) { + MasterMetricManager::instance().dec_soft_pin_key_count(-metric_delta); + } +} + +size_t MasterService::GetMetadataShardIndex( + const ObjectMetadata& metadata) const { + return metadata.group_id.empty() + ? getShardIndex(metadata.tenant_id, metadata.user_key) + : getShardIndex(metadata.group_id); +} + +void MasterService::ApplySoftPinEvaluation( + const ObjectMetadata& metadata, + const ObjectMetadata::SoftPinEvaluation& result) const { + if (!result.deadline_to_index && !result.removed_deadline) { + ApplySoftPinMetricDelta(result.metric_delta); + return; + } + + const size_t shard_idx = GetMetadataShardIndex(metadata); + const auto scoped_key = metadata.tenant_id.MakeScopedKey(metadata.user_key); + if (result.deadline_to_index) { + soft_pin_deadline_index_.Upsert(scoped_key, shard_idx, + *result.deadline_to_index); + } else if (result.removed_deadline) { + soft_pin_deadline_index_.RemoveIfMatches(scoped_key, shard_idx, + *result.removed_deadline); + } + ApplySoftPinMetricDelta(result.metric_delta); +} + +bool MasterService::IsSoftPinActive( + const ObjectMetadata& metadata, + const std::chrono::system_clock::time_point& now) const { + const auto evaluation = metadata.EvaluateSoftPin(now); + ApplySoftPinEvaluation(metadata, evaluation); + return evaluation.active; +} + +void MasterService::CleanupExpiredSoftPins( + const std::chrono::system_clock::time_point& now) { + auto expired_entries = soft_pin_deadline_index_.PopExpired(now); + std::sort(expired_entries.begin(), expired_entries.end(), + [](const auto& lhs, const auto& rhs) { + return lhs.shard_idx < rhs.shard_idx; + }); + + int expired_count = 0; + for (size_t begin = 0; begin < expired_entries.size();) { + const size_t shard_idx = expired_entries[begin].shard_idx; + size_t end = begin + 1; + while (end < expired_entries.size() && + expired_entries[end].shard_idx == shard_idx) { + ++end; + } + + MetadataShardAccessorRO shard(this, shard_idx); + for (size_t i = begin; i < end; ++i) { + const auto& entry = expired_entries[i]; + const auto [tenant_id, key] = + TenantId::ParseScopedKey(entry.scoped_key); + const auto tenant_it = shard->tenants.find(tenant_id); + if (tenant_it == shard->tenants.end()) { + continue; + } + const auto metadata_it = tenant_it->second.metadata.find(key); + if (metadata_it == tenant_it->second.metadata.end()) { + continue; + } + if (metadata_it->second.ExpireSoftPinIfDeadlineMatches( + entry.deadline, now)) { + ++expired_count; + } + } + begin = end; + } + if (expired_count > 0) { + MasterMetricManager::instance().dec_soft_pin_key_count(expired_count); + } +} + void MasterService::GrantLeaseForGroup(const TenantState& tenant_state, const std::string& key, const ObjectMetadata& metadata) const { if (!metadata.IsGrouped()) { - metadata.GrantLease(default_kv_lease_ttl_, default_kv_soft_pin_ttl_); + metadata.GrantReadLease(default_kv_lease_ttl_); return; } - bool needs_refresh = metadata.NeedsLeaseRefresh(default_kv_lease_ttl_, - default_kv_soft_pin_ttl_); + bool needs_refresh = metadata.NeedsReadLeaseRefresh(default_kv_lease_ttl_); if (!needs_refresh) { std::shared_lock lock(group_routing_mutex_); needs_refresh = @@ -2022,19 +2241,18 @@ void MasterService::GrantLeaseForGroup(const TenantState& tenant_state, auto group_it = tenant_state.group_members.find(metadata.group_id); if (group_it == tenant_state.group_members.end()) { - metadata.GrantLease(default_kv_lease_ttl_, default_kv_soft_pin_ttl_); + metadata.GrantReadLease(default_kv_lease_ttl_); return; } for (const auto& member_key : group_it->second) { auto mit = tenant_state.metadata.find(member_key); if (mit != tenant_state.metadata.end()) { - mit->second.GrantLease(default_kv_lease_ttl_, - default_kv_soft_pin_ttl_); + mit->second.GrantReadLease(default_kv_lease_ttl_); } } if (group_it->second.find(key) == group_it->second.end()) { - metadata.GrantLease(default_kv_lease_ttl_, default_kv_soft_pin_ttl_); + metadata.GrantReadLease(default_kv_lease_ttl_); } { std::unique_lock lock(group_routing_mutex_); @@ -2136,9 +2354,12 @@ void MasterService::TaskCleanupThreadFunc() { } std::shared_lock shared_lock(snapshot_mutex_); - auto write_access = task_manager_.get_write_access(); - write_access.prune_expired_tasks(); - write_access.prune_finished_tasks(); + { + auto write_access = task_manager_.get_write_access(); + write_access.prune_expired_tasks(); + write_access.prune_finished_tasks(); + } + CleanupExpiredSoftPins(std::chrono::system_clock::now()); } LOG(INFO) << "Task cleanup thread stopped"; } @@ -2301,7 +2522,7 @@ auto MasterService::ExistKey(const std::string& key, const TenantId& tenant_id) if (ts) { GrantLeaseForGroup(*ts, key, metadata); } else { - metadata.GrantLease(default_kv_lease_ttl_, default_kv_soft_pin_ttl_); + metadata.GrantReadLease(default_kv_lease_ttl_); } return true; } @@ -2617,8 +2838,9 @@ void MasterService::RestoreFromStandbySnapshot( std::piecewise_construct, std::forward_as_tuple(user_key), std::forward_as_tuple( standby_meta.client_id, now, standby_meta.size, - std::move(replicas), false, false, standby_meta.data_type, - standby_meta.group_id, tenant_id, user_key)); + std::move(replicas), std::nullopt, false, + standby_meta.data_type, standby_meta.group_id, tenant_id, + user_key)); if (!standby_meta.group_id.empty()) { RegisterGroupMember(tenant_state, tenant_id, user_key, standby_meta.group_id); @@ -3042,8 +3264,7 @@ auto MasterService::GetReplicaList(const std::string& key, if (ts) { GrantLeaseForGroup(*ts, key, metadata); } else { - metadata.GrantLease(default_kv_lease_ttl_, - default_kv_soft_pin_ttl_); + metadata.GrantReadLease(default_kv_lease_ttl_); } // Promotion-on-hit eligibility: only when no MEMORY replica is @@ -3327,8 +3548,12 @@ auto MasterService::AllocateAndInsertMetadata( MetadataShardAccessorRW& shard, const UUID& client_id, const std::string& key, uint64_t value_length, const ReplicateConfig& config, const std::string& group_id, - const TenantId& tenant_id, const std::chrono::system_clock::time_point& now) + const TenantId& tenant_id, const std::chrono::system_clock::time_point& now, + const ResolvedSoftPinRequest& soft_pin_request, + std::optional + committed_soft_pin_timeout) -> tl::expected, ErrorCode> { + const auto deadline_to_index = committed_soft_pin_timeout; auto& tenant_state = shard->tenants[tenant_id]; if (tenant_state.metadata.contains(key)) { LOG(INFO) << "key=" << key << ", info=object_already_exists"; @@ -3513,13 +3738,16 @@ auto MasterService::AllocateAndInsertMetadata( } std::vector replica_list; + std::vector eligible_replica_ids; replica_list.reserve(replicas.size()); + eligible_replica_ids.reserve(replicas.size()); int i = 0; VLOG(1) << "PutStart, create replicas: client_id=" << client_id << ", key=" << key << ", value_length=" << value_length; for (const auto& replica : replicas) { const auto desc = replica.get_descriptor(); replica_list.emplace_back(desc); + eligible_replica_ids.push_back(replica.id()); if (replica.is_memory_replica()) { const auto& mem_desc = desc.get_memory_descriptor(); @@ -3539,14 +3767,22 @@ auto MasterService::AllocateAndInsertMetadata( auto [it, inserted] = tenant_state.metadata.emplace( std::piecewise_construct, std::forward_as_tuple(key), std::forward_as_tuple(client_id, now, value_length, std::move(replicas), - config.with_soft_pin, config.with_hard_pin, - config.data_type, group_id, tenant_id, key)); + std::move(committed_soft_pin_timeout), + config.with_hard_pin, config.data_type, group_id, + tenant_id, key)); if (!inserted) { LOG(INFO) << "key=" << key << ", info=object_already_exists"; abort_reserved_quota(); return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); } IncrementTenantMetadataObjectCount(tenant_id); + it->second.BeginSoftPinAction(soft_pin_request, + std::move(eligible_replica_ids)); + if (deadline_to_index) { + soft_pin_deadline_index_.Upsert(tenant_id.MakeScopedKey(key), + GetMetadataShardIndex(it->second), + *deadline_to_index); + } it->second.reserved_quota_charge_bytes = reserved_quota_charge; RegisterGroupMember(tenant_state, tenant_id, key, group_id); tenant_state.processing_keys.insert(key); @@ -3590,6 +3826,11 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, } #endif + auto soft_pin_request = ResolveSoftPinRequest(config); + if (!soft_pin_request) { + return tl::make_unexpected(soft_pin_request.error()); + } + UpdateClientHostId(client_id, config.host_id); if ((memory_allocator_type_ == BufferAllocatorType::CACHELIB) && @@ -3705,7 +3946,7 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, } else { return AllocateAndInsertMetadata( shard, client_id, key, slice_length, config, group_id, - object_id.tenant_id, now); + object_id.tenant_id, now, *soft_pin_request); } } } @@ -3721,7 +3962,7 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, } return AllocateAndInsertMetadata(shard, client_id, key, slice_length, config, group_id, object_id.tenant_id, - now); + now, *soft_pin_request); }; for (int attempt = 0; attempt <= kMaxTenantQuotaEvictionRetries; @@ -3805,11 +4046,27 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, return tl::make_unexpected(ErrorCode::INVALID_WRITE); } + const bool had_completed_replica = + metadata.HasReplica(&Replica::fn_is_completed); + bool completed_pending_replica = false; metadata.VisitReplicas( [&is_target_replica](const Replica& replica) { return replica.is_processing() && is_target_replica(replica); }, - [](Replica& replica) { replica.mark_complete(); }); + [&metadata, &completed_pending_replica](Replica& replica) { + if (replica.is_processing() && + metadata.PendingSoftPinOwnsReplica(replica.id())) { + completed_pending_replica = true; + } + replica.mark_complete(); + }); + + if (!had_completed_replica && completed_pending_replica && + metadata.HasReplica(&Replica::fn_is_completed)) { + const auto soft_pin_result = + metadata.CommitPendingSoftPin(std::chrono::system_clock::now()); + ApplySoftPinEvaluation(metadata, soft_pin_result); + } if (object_meta.object_checksum.has_value() || replica_type == ReplicaType::ALL || @@ -3872,10 +4129,7 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, SyncCacheTotalAccounting(metadata); // TODO: add inc_nof_cache_nums() (ranhaojia) - // 1. Set lease timeout to now, indicating that the object has no lease - // at beginning. 2. If this object has soft pin enabled, set it to be soft - // pinned. - metadata.GrantLease(0, default_kv_soft_pin_ttl_); + metadata.GrantReadLease(0); PublishKvStored(key, replica_type, metadata, object_id.tenant_id); if (enable_oplog_ && ordered_oplog_writer_) { @@ -3913,7 +4167,7 @@ auto MasterService::AddReplica(const UUID& client_id, const std::string& key, accessor.Create( client_id, replica.get_descriptor().get_local_disk_descriptor().object_size, - std::vector{}, false); + std::vector{}); } auto& metadata = accessor.Get(); if (replica.type() != ReplicaType::LOCAL_DISK) { @@ -4054,6 +4308,7 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, removed_ids.push_back(r.id()); r.mark_removed(); }); + metadata.ClearPendingSoftPinIfNoViableReplica(); tl::expected persist_result; if (remaining.empty()) { @@ -4086,6 +4341,7 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, const uint64_t before_charge = CompletedMemoryQuotaCharge(metadata); EraseReplicasWithCacheTotalAccounting(metadata, target_pred); + metadata.ClearPendingSoftPinIfNoViableReplica(); const uint64_t after_charge = CompletedMemoryQuotaCharge(metadata); if (before_charge > after_charge) { ReleaseCommittedQuotaCharge(metadata, before_charge - after_charge); @@ -4192,6 +4448,11 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, } #endif + auto soft_pin_request = ResolveSoftPinRequest(config); + if (!soft_pin_request) { + return tl::make_unexpected(soft_pin_request.error()); + } + UpdateClientHostId(client_id, config.host_id); if ((memory_allocator_type_ == BufferAllocatorType::CACHELIB) && @@ -4231,6 +4492,8 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, auto now = std::chrono::system_clock::now(); std::optional case_a_retry_shard_idx; + std::optional + case_a_committed_soft_pin_timeout; { // --- Lock acquisition --- auto alive_clients = getAliveClientsSnapshot(); @@ -4317,6 +4580,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, if (tenant_state.processing_keys.count(key) > 0) { auto processing_replicas = metadata.PopReplicas(&Replica::fn_is_processing); + metadata.ClearPendingSoftPinAction(); if (!processing_replicas.empty()) { std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back( @@ -4328,6 +4592,12 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, // If no COMPLETE replicas survive the preemption, this key // effectively does not exist — fall through to Case A. if (!metadata.HasReplica(&Replica::fn_is_completed)) { + case_a_committed_soft_pin_timeout = + metadata.GetCommittedSoftPinTimeout(); + if (case_a_committed_soft_pin_timeout && + *case_a_committed_soft_pin_timeout <= now) { + case_a_committed_soft_pin_timeout.reset(); + } EraseMetadata(tenant_state, it, object_id.tenant_id, QuotaEraseMode::kFull, &shard); it = tenant_state.metadata.end(); @@ -4351,7 +4621,8 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, } else { return AllocateAndInsertMetadata( shard, client_id, key, slice_length, config, group_id, - object_id.tenant_id, now); + object_id.tenant_id, now, *soft_pin_request, + std::move(case_a_committed_soft_pin_timeout)); } } else { // --- Step 2: key exists with COMPLETE replicas → Case B or C @@ -4379,28 +4650,18 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, metadata.client_id = client_id; metadata.put_start_time = now; - // Reconcile soft_pin state with the incoming config. - { - SpinLocker locker(&metadata.lock); - if (config.with_soft_pin && - !metadata.soft_pin_timeout) { - metadata.soft_pin_timeout.emplace(); - MasterMetricManager::instance() - .inc_soft_pin_key_count(1); - } else if (!config.with_soft_pin && - metadata.soft_pin_timeout) { - metadata.soft_pin_timeout.reset(); - MasterMetricManager::instance() - .dec_soft_pin_key_count(1); - } - } - // Mark COMPLETE → PROCESSING so readers won't see stale // data mid-transfer. The key becomes unreadable until // UpsertEnd. + std::vector eligible_replica_ids; metadata.VisitReplicas( &Replica::fn_is_completed, - [](Replica& replica) { replica.mark_processing(); }); + [&eligible_replica_ids](Replica& replica) { + eligible_replica_ids.push_back(replica.id()); + replica.mark_processing(); + }); + metadata.BeginSoftPinAction( + *soft_pin_request, std::move(eligible_replica_ids)); SyncCacheTotalAccounting(metadata); tenant_state.processing_keys.insert(key); @@ -4432,8 +4693,12 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, ReplicateConfig merged_config = config; merged_config.with_hard_pin = merged_config.with_hard_pin || metadata.IsHardPinned(); - merged_config.with_soft_pin = - merged_config.with_soft_pin || metadata.IsSoftPinned(); + auto committed_soft_pin_timeout = + metadata.GetCommittedSoftPinTimeout(); + if (committed_soft_pin_timeout && + *committed_soft_pin_timeout <= now) { + committed_soft_pin_timeout.reset(); + } const std::string existing_group_id = metadata.group_id; const uint64_t old_quota_charge = @@ -4455,7 +4720,8 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, << ", action=upsert_start_case_c_reallocate"; auto allocate_result = AllocateAndInsertMetadata( shard, client_id, key, slice_length, merged_config, - existing_group_id, object_id.tenant_id, now); + existing_group_id, object_id.tenant_id, now, + *soft_pin_request, std::move(committed_soft_pin_timeout)); if (!allocate_result) { ReleaseTenantQuota(object_id.tenant_id, old_quota_charge); return allocate_result; @@ -4478,9 +4744,10 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, LOG(INFO) << "key=" << key << ", info=object_already_exists"; return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); } - return AllocateAndInsertMetadata(shard, client_id, key, slice_length, - config, group_id, object_id.tenant_id, - now); + return AllocateAndInsertMetadata( + shard, client_id, key, slice_length, config, group_id, + object_id.tenant_id, now, *soft_pin_request, + std::move(case_a_committed_soft_pin_timeout)); }; for (int attempt = 0; attempt <= kMaxTenantQuotaEvictionRetries; @@ -7298,6 +7565,7 @@ void MasterService::DiscardExpiredProcessingReplicas( auto& metadata = it->second; if (!metadata.IsValid() || metadata.AllReplicas(&Replica::fn_is_completed)) { + metadata.ClearPendingSoftPinIfNoViableReplica(); if (!metadata.IsValid()) { auto next_key_it = std::next(key_it); EraseMetadata(tenant_state, it, tenant_it->first, @@ -7365,6 +7633,7 @@ void MasterService::DiscardExpiredProcessingReplicas( // Persist OK (or HA disabled / never published) — apply. auto replicas = metadata.PopReplicas(&Replica::fn_is_processing); + metadata.ClearPendingSoftPinIfNoViableReplica(); if (!replicas.empty()) { discarded_replicas.emplace_back(std::move(replicas), ttl); } @@ -7720,8 +7989,7 @@ tl::expected MasterService::ApplySnapshotState( it != tenant_state.metadata.end();) { if (it->second.HasDiffRepStatus( ReplicaStatus::COMPLETE) || - (it->second.IsLeaseExpired(cleanup_now) && - !it->second.IsSoftPinned(cleanup_now))) { + it->second.IsLeaseExpired(cleanup_now)) { VLOG(1) << "clear metadata key=" << it->first; it = EraseMetadata(tenant_state, it, tenant_it->first); @@ -7777,6 +8045,9 @@ tl::expected MasterService::ApplySnapshotState( << MasterMetricManager::instance().get_allocated_mem_size(); } + // Soft pin is runtime-only and is never restored from a snapshot. + soft_pin_deadline_index_.Clear(); + // Rebuild total capacity metrics { MasterMetricManager::instance().reset_total_mem_capacity(); @@ -7953,7 +8224,7 @@ MasterService::EvictTenantMemoryForQuota(const TenantId& tenant_id, auto& member_metadata = member_it->second; if (member_metadata.IsHardPinned() || !member_metadata.IsLeaseExpired(now) || - (!allow_soft_pinned && member_metadata.IsSoftPinned(now)) || + (!allow_soft_pinned && IsSoftPinActive(member_metadata, now)) || !can_evict_replicas(member_metadata)) { continue; } @@ -7991,7 +8262,8 @@ MasterService::EvictTenantMemoryForQuota(const TenantId& tenant_id, auto& metadata = it->second; if (metadata.IsHardPinned() || !metadata.IsLeaseExpired(now) || - (!allow_soft_pinned && metadata.IsSoftPinned(now)) || + (!allow_soft_pinned && + IsSoftPinActive(metadata, now)) || !can_evict_replicas(metadata)) { ++it; continue; @@ -8315,7 +8587,7 @@ void MasterService::BatchEvict(double evict_ratio_target, auto& member_metadata = member_it->second; if (member_metadata.IsHardPinned() || !member_metadata.IsLeaseExpired(now) || - (!allow_soft_pinned && member_metadata.IsSoftPinned(now)) || + (!allow_soft_pinned && IsSoftPinActive(member_metadata, now)) || !can_evict_replicas(member_metadata)) { continue; } @@ -8408,7 +8680,7 @@ void MasterService::BatchEvict(double evict_ratio_target, if (has_evictable) shard_evictable_count++; if (!it->second.IsLeaseExpired(now) || !has_evictable) continue; - if (!it->second.IsSoftPinned(now)) { + if (!IsSoftPinActive(it->second, now)) { if (compact_frontier_prebypass) { local_candidates[t].push_back( {s, tenant_id, it->first, @@ -8510,7 +8782,7 @@ void MasterService::BatchEvict(double evict_ratio_target, tenant_state.metadata) { if (metadata.IsHardPinned() || !metadata.IsLeaseExpired(now) || - metadata.IsSoftPinned(now) || + IsSoftPinActive(metadata, now) || !can_evict_replicas(metadata)) { continue; } @@ -8620,7 +8892,7 @@ void MasterService::BatchEvict(double evict_ratio_target, // Re-validate: state may have changed since Phase 1 if (!it->second.IsLeaseExpired(now) || - it->second.IsSoftPinned(now) || + IsSoftPinActive(it->second, now) || !can_evict_replicas(it->second)) { no_pin_objects.push_back(c.lease_timeout); continue; @@ -8709,7 +8981,7 @@ void MasterService::BatchEvict(double evict_ratio_target, if (!it->second.IsHardPinned() && it->second.IsLeaseExpired(now) && it->second.lease_timeout <= target_timeout && - !it->second.IsSoftPinned(now) && + !IsSoftPinActive(it->second, now) && can_evict_replicas(it->second)) { auto evict_result = try_evict_group_or_object( tenant_it->first, it->first, it->second, @@ -8773,7 +9045,7 @@ void MasterService::BatchEvict(double evict_ratio_target, ++it; continue; } - if (!it->second.IsSoftPinned(now) || + if (!IsSoftPinActive(it->second, now) || it->second.lease_timeout <= soft_target_timeout) { auto evict_result = try_evict_group_or_object( @@ -8932,7 +9204,7 @@ void MasterService::NoFBatchEvict(double evict_ratio_target, shard_evicted_count < ideal_evict_num;) { auto& metadata = it->second; if (metadata.IsHardPinned() || !metadata.IsLeaseExpired(now) || - metadata.IsSoftPinned(now)) { + IsSoftPinActive(metadata, now)) { ++it; continue; } @@ -9659,6 +9931,7 @@ MasterService::MetadataSerializer::Deserialize( } void MasterService::MetadataSerializer::Reset() { + service_->soft_pin_deadline_index_.Clear(); for (auto& shard : service_->metadata_shards_) { shard.tenants.clear(); } @@ -9801,13 +10074,11 @@ MasterService::MetadataSerializer::DeserializeShard(const msgpack::object& obj, std::piecewise_construct, std::forward_as_tuple(std::move(key)), std::forward_as_tuple( metadata_ptr->client_id, metadata_ptr->put_start_time, - metadata_ptr->size, metadata_ptr->PopReplicas(), - metadata_ptr->soft_pin_timeout.has_value(), + metadata_ptr->size, metadata_ptr->PopReplicas(), std::nullopt, metadata_ptr->IsHardPinned(), metadata_ptr->data_type, metadata_ptr->group_id, tenant_id, user_key)); it->second.lease_timeout = metadata_ptr->lease_timeout; - it->second.soft_pin_timeout = metadata_ptr->soft_pin_timeout; it->second.object_checksum = metadata_ptr->object_checksum; // Recompute disk_object_count for restored metadata @@ -9859,18 +10130,10 @@ MasterService::MetadataSerializer::SerializeMetadata( .count(); packer.pack(lease_timestamp); - // Serialize soft_pin_timeout (if exists) - if (metadata.soft_pin_timeout.has_value()) { - packer.pack(true); // Mark soft_pin_timeout exists - auto soft_pin_timestamp = - std::chrono::duration_cast( - metadata.soft_pin_timeout.value().time_since_epoch()) - .count(); - packer.pack(soft_pin_timestamp); - } else { - packer.pack(false); // Mark soft_pin_timeout does not exist - packer.pack(uint64_t(0)); // Placeholder - } + // Keep the legacy snapshot slots for format compatibility, but soft pin is + // runtime-only state and is intentionally not persisted. + packer.pack(false); + packer.pack(uint64_t(0)); // Serialize replicas count packer.pack(static_cast(metadata.CountReplicas())); @@ -9936,11 +10199,22 @@ MasterService::MetadataSerializer::DeserializeMetadata( // Deserialize lease_timeout uint64_t lease_timestamp = array[index++].as(); - // Deserialize soft_pin_timeout flag - bool has_soft_pin_timeout = array[index++].as(); + // Parse and discard the legacy soft-pin fields. Recovered objects always + // become ordinary cache. + (void)array[index++].as(); + (void)array[index++].as(); - // Deserialize soft_pin_timeout value - uint64_t soft_pin_timestamp = array[index++].as(); + const auto max_timestamp = + std::chrono::duration_cast( + std::chrono::system_clock::time_point::max().time_since_epoch()) + .count(); + if (max_timestamp < 0 || + put_start_time_timestamp > static_cast(max_timestamp) || + lease_timestamp > static_cast(max_timestamp)) { + return tl::unexpected(SerializationError( + ErrorCode::DESERIALIZE_FAIL, + "ObjectMetadata timestamp exceeds system_clock range")); + } // Deserialize replicas count uint32_t replicas_count = array[index++].as(); @@ -10016,25 +10290,17 @@ MasterService::MetadataSerializer::DeserializeMetadata( "deserialize ObjectMetadata optional field type mismatch")); } - // Create ObjectMetadata instance - bool enable_soft_pin = has_soft_pin_timeout; + // Create ObjectMetadata instance. Soft pin is not restored. auto metadata = std::make_unique( client_id, std::chrono::system_clock::time_point( std::chrono::milliseconds(put_start_time_timestamp)), - size, std::move(replicas), enable_soft_pin, is_hard_pinned, data_type, + size, std::move(replicas), std::nullopt, is_hard_pinned, data_type, group_id); metadata->object_checksum = object_checksum; metadata->lease_timeout = std::chrono::system_clock::time_point( std::chrono::milliseconds(lease_timestamp)); - // Set soft_pin_timeout (if exists) - if (has_soft_pin_timeout) { - metadata->soft_pin_timeout.emplace( - std::chrono::system_clock::time_point( - std::chrono::milliseconds(soft_pin_timestamp))); - } - return metadata; } diff --git a/mooncake-store/src/store_c.cpp b/mooncake-store/src/store_c.cpp index a3a19086ef..09084f52a8 100644 --- a/mooncake-store/src/store_c.cpp +++ b/mooncake-store/src/store_c.cpp @@ -44,7 +44,9 @@ mooncake::ReplicateConfig to_replicate_config( if (!c_config) return config; config.replica_num = c_config->replica_num; - config.with_soft_pin = c_config->with_soft_pin != 0; + config.soft_pin_action = c_config->with_soft_pin != 0 + ? mooncake::SoftPinAction::ENABLE + : mooncake::SoftPinAction::PRESERVE; config.with_hard_pin = c_config->with_hard_pin != 0; if (c_config->preferred_segments && c_config->preferred_segments_count > 0) { diff --git a/mooncake-store/tests/batch_evict_test.cpp b/mooncake-store/tests/batch_evict_test.cpp index 9086b63fc1..84ad36a749 100644 --- a/mooncake-store/tests/batch_evict_test.cpp +++ b/mooncake-store/tests/batch_evict_test.cpp @@ -77,7 +77,8 @@ class BatchEvictTest : public ::testing::Test { ReplicateConfig config; config.replica_num = 1; config.preferred_segment = kSegmentName; - config.with_soft_pin = with_soft_pin; + config.soft_pin_action = + with_soft_pin ? SoftPinAction::ENABLE : SoftPinAction::PRESERVE; if (!group_id.empty()) { config.group_ids = std::vector{group_id}; } diff --git a/mooncake-store/tests/ha/snapshot/catalog_backed_snapshot_provider_test.cpp b/mooncake-store/tests/ha/snapshot/catalog_backed_snapshot_provider_test.cpp index c81d5eccf8..533a524086 100644 --- a/mooncake-store/tests/ha/snapshot/catalog_backed_snapshot_provider_test.cpp +++ b/mooncake-store/tests/ha/snapshot/catalog_backed_snapshot_provider_test.cpp @@ -2,11 +2,11 @@ #include #include +#include #include #include #include #include -#include #include #include @@ -203,6 +203,30 @@ TEST_P(CatalogBackedSnapshotProviderTest, ExpectLoadsDefaultObject(); } +TEST_P(CatalogBackedSnapshotProviderTest, + LoadLatestSnapshotIgnoresSoftPinForRetention) { + const auto now_ms = std::chrono::duration_cast( + std::chrono::system_clock::now().time_since_epoch()) + .count(); + ASSERT_GT(now_ms, 0); + auto published = mooncake::test::PublishSnapshotPayloadBytes( + *object_store_, *catalog_store_, descriptor_, + BuildMetadataPayload( + UUID{1, 2}, kDefaultTestObjectKey, kDefaultTestDiskFilePath, + kDefaultTestObjectSize, kDefaultTestPutStartTimeMs, + static_cast(now_ms - 1), SnapshotMetadataFormat::kLegacy, + static_cast(now_ms + 60'000))); + ASSERT_TRUE(published.has_value()) << published.error(); + snapshot_published_ = true; + + auto provider = CreateProvider(); + ASSERT_TRUE(provider.has_value()) << toString(provider.error()); + auto snapshot = provider.value()->LoadLatestSnapshot(cluster_id_); + ASSERT_TRUE(snapshot.has_value()) << toString(snapshot.error()); + ASSERT_TRUE(snapshot->has_value()); + EXPECT_TRUE(snapshot->value().metadata.empty()); +} + TEST_P(CatalogBackedSnapshotProviderTest, RejectsOverflowingReplicaCount) { // A near-UINT32_MAX replica_count must not wrap the format-detection // arithmetic into a valid-looking total and slip an out-of-bounds index diff --git a/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot.cpp b/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot.cpp index 4231a9ef99..c129c141ed 100644 --- a/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot.cpp +++ b/mooncake-store/tests/ha/snapshot/master_service_test_for_snapshot.cpp @@ -1642,7 +1642,7 @@ TEST_F(MasterServiceSnapshotTest, RemoveSoftPinObject) { uint64_t slice_length = 1024; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; // Verify soft pin does not block remove ASSERT_TRUE(service_ @@ -1698,7 +1698,7 @@ TEST_F(MasterServiceSnapshotTest, SoftPinObjectsNotEvictedBeforeOtherObjects) { uint64_t slice_length = value_size; ReplicateConfig soft_pin_config; soft_pin_config.replica_num = 1; - soft_pin_config.with_soft_pin = true; + soft_pin_config.soft_pin_action = SoftPinAction::ENABLE; ASSERT_TRUE(service_ ->PutStart(client_id, pin_key, TenantId::Default(), @@ -1779,7 +1779,7 @@ TEST_F(MasterServiceSnapshotTest, SoftPinObjectsCanBeEvicted) { uint64_t slice_length = value_size; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; if (service_ ->PutStart(client_id, key, TenantId::Default(), slice_length, config) @@ -1807,101 +1807,53 @@ TEST_F(MasterServiceSnapshotTest, SoftPinObjectsCanBeEvicted) { // service_->RemoveAll(); } -TEST_F(MasterServiceSnapshotTest, SoftPinExtendedOnGet) { +TEST_F(MasterServiceSnapshotTest, SoftPinExpiresAndGetDoesNotReactivate) { const uint64_t kv_lease_ttl = 200; - // The soft pin ttl shall not be too large, otherwise the test will take too - // long - const uint64_t kv_soft_pin_ttl = 1000; - static_assert( - kv_soft_pin_ttl > kv_lease_ttl, - "kv_soft_pin_ttl must be larger than kv_lease_ttl in this test"); - const double eviction_ratio = 0.5; - const bool allow_evict_soft_pinned_objects = true; + const uint64_t kv_soft_pin_ttl = 20; auto service_config = MasterServiceConfig::builder() .set_default_kv_lease_ttl(kv_lease_ttl) .set_default_kv_soft_pin_ttl(kv_soft_pin_ttl) - .set_allow_evict_soft_pinned_objects( - allow_evict_soft_pinned_objects) - .set_eviction_ratio(eviction_ratio) .build(); service_.reset(new MasterService(service_config)); const UUID client_id = generate_uuid(); - // Mount segment and put an object constexpr size_t buffer = 0x300000000; constexpr size_t segment_size = 1024 * 1024 * 16; - constexpr size_t value_size = 1024 * 1024; + constexpr size_t value_size = 1024; [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_, "test_segment", buffer, segment_size); - // The eviction has random factors, so test 3 times - for (int test_i = 0; test_i < 3; test_i++) { - // Put pin_key first - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - uint64_t slice_length = value_size; - ReplicateConfig soft_pin_config; - soft_pin_config.replica_num = 1; - soft_pin_config.with_soft_pin = true; - - ASSERT_TRUE(service_->PutStart(client_id, pin_key, - TenantId::Default(), slice_length, - soft_pin_config)); - ASSERT_TRUE(service_ - ->PutEnd(client_id, pin_key, TenantId::Default(), - ReplicaType::MEMORY) - .has_value()); - } - - // Wait for the soft pin to expire - std::this_thread::sleep_for(std::chrono::milliseconds(kv_soft_pin_ttl)); - - // Get the pin_key to extend the soft pin - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - ASSERT_TRUE(service_->GetReplicaList(pin_key, TenantId::Default()) - .has_value()); - } - - // Fill the segment to trigger eviction - int failed_puts = 0; - for (int i = 0; i < 16; i++) { - std::string key = "key" + std::to_string(i); - uint64_t slice_length = value_size; - ReplicateConfig config; - config.replica_num = 1; - if (service_ - ->PutStart(client_id, key, TenantId::Default(), - slice_length, config) - .has_value()) { - ASSERT_TRUE(service_ - ->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY) - .has_value()); - } else { - failed_puts++; - } - } - ASSERT_GT(failed_puts, 0); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + ReplicateConfig config; + config.soft_pin_action = SoftPinAction::ENABLE; + ASSERT_TRUE(service_ + ->PutStart(client_id, "pin_key", TenantId::Default(), + value_size, config) + .has_value()); + ASSERT_TRUE(service_ + ->PutEnd(client_id, "pin_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); - // wait for eviction - std::this_thread::sleep_for(std::chrono::milliseconds(kv_lease_ttl)); + std::this_thread::sleep_for( + std::chrono::milliseconds(kv_soft_pin_ttl + 10)); + ASSERT_TRUE( + service_->GetReplicaList("pin_key", TenantId::Default()).has_value()); - // pin_key should still be accessible - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - ASSERT_TRUE(service_->GetReplicaList(pin_key, TenantId::Default()) - .has_value()); - } - // [Commented for snapshot test] The following RemoveAll would clear - // data before TearDown snapshot verification Only remove all objects - // before the next turn (skip last round) - if (test_i < 2) { - std::this_thread::sleep_for( - std::chrono::milliseconds(kv_lease_ttl)); - service_->RemoveAll(); - } - } + ReplicateConfig preserve; + ASSERT_TRUE(service_ + ->UpsertStart(client_id, "pin_key", TenantId::Default(), + value_size, preserve) + .has_value()); + ASSERT_TRUE(service_ + ->UpsertEnd(client_id, "pin_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); } TEST_F(MasterServiceSnapshotTest, SoftPinObjectsNotAllowEvict) { @@ -1934,7 +1886,7 @@ TEST_F(MasterServiceSnapshotTest, SoftPinObjectsNotAllowEvict) { uint64_t slice_length = value_size; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; if (service_ ->PutStart(client_id, key, TenantId::Default(), slice_length, config) @@ -2656,7 +2608,7 @@ TEST_F(MasterServiceSnapshotTest, BatchReplicaClearSpecificSegment) { ASSERT_TRUE(put_end_result.has_value()); // 4. Wait for lease to expire and verify it's actually expired - // PutEnd calls GrantLease(0, ...) which sets lease_timeout to now. + // PutEnd grants a zero-duration read lease, setting lease_timeout to now. // Due to clock precision and timing, we need to ensure the lease is // actually expired before calling BatchReplicaClear. // Use a small delay and then poll to ensure lease is expired. diff --git a/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp b/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp index 62e818000f..2b4ba1821f 100644 --- a/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp +++ b/mooncake-store/tests/ha/snapshot/snapshot_child_process_test.cpp @@ -18,6 +18,7 @@ #include #include #include +#include #include #include #include @@ -239,6 +240,26 @@ class SnapshotChildProcessTest : public ::testing::Test { tenant_it->second.metadata.end(); } + size_t SoftPinRegistrationCount(MasterService* svc) { + return svc->soft_pin_deadline_index_.RegistrationCountForTest(); + } + + std::optional GetSoftPinDeadline( + MasterService* svc, const std::string& key) { + const size_t shard_idx = + svc->getMetadataShardIndex(TenantId::Default(), key); + MasterService::MetadataShardAccessorRO shard(svc, shard_idx); + const auto tenant_it = shard->tenants.find(TenantId::Default()); + if (tenant_it == shard->tenants.end()) { + return std::nullopt; + } + const auto metadata_it = tenant_it->second.metadata.find(key); + if (metadata_it == tenant_it->second.metadata.end()) { + return std::nullopt; + } + return metadata_it->second.GetCommittedSoftPinTimeout(); + } + uint32_t GetShardIndexForTest(const std::string& key) { return static_cast(service_->getShardIndex(key)); } @@ -1154,7 +1175,8 @@ TEST_F(SnapshotChildProcessTest, RestoreCleansExpiredLease) { ASSERT_TRUE(mount_result.has_value()) << "MountSegment failed"; // Add two complete objects via PutStart + PutEnd - // Note: PutEnd calls GrantLease(0, ...) so lease is immediately expired + // PutEnd only grants a zero-duration read lease, so it is immediately + // expired. std::string expired_key = "expired_lease_object"; auto put_exp = service_->PutStart(client_id, expired_key, TenantId::Default(), {1024}, @@ -1176,9 +1198,27 @@ TEST_F(SnapshotChildProcessTest, RestoreCleansExpiredLease) { .has_value()) << "PutEnd normal failed"; + const int64_t soft_pin_baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + std::string soft_pinned_key = "soft_pin_valid_lease_object"; + ReplicateConfig soft_pin_config; + soft_pin_config.soft_pin_action = SoftPinAction::ENABLE; + soft_pin_config.soft_pin_ttl_ms = 60'000; + auto put_soft = + service_->PutStart(client_id, soft_pinned_key, TenantId::Default(), + {1024}, soft_pin_config); + ASSERT_TRUE(put_soft.has_value()) << "PutStart soft pin failed"; + ASSERT_TRUE(service_ + ->PutEnd(client_id, soft_pinned_key, TenantId::Default(), + ReplicaType::MEMORY) + .has_value()) + << "PutEnd soft pin failed"; + // ExistKey grants a fresh lease (now + 600s) to normal_key EXPECT_TRUE( service_->ExistKey(normal_key, TenantId::Default()).value_or(false)); + EXPECT_TRUE(service_->ExistKey(soft_pinned_key, TenantId::Default()) + .value_or(false)); // Do NOT call ExistKey on expired_key, its lease stays expired from PutEnd // Step 2: Persist state @@ -1205,6 +1245,17 @@ TEST_F(SnapshotChildProcessTest, RestoreCleansExpiredLease) { << "Normal object with valid lease should survive restore"; EXPECT_FALSE(KeyExistsInMetadata(restored_service.get(), expired_key)) << "Lease-expired object should be cleaned during restore"; + EXPECT_TRUE(restored_service->ExistKey(soft_pinned_key, TenantId::Default()) + .value_or(false)) + << "Soft-pinned object with a valid read lease should restore as cache"; + EXPECT_FALSE( + GetSoftPinDeadline(restored_service.get(), soft_pinned_key).has_value()) + << "Snapshot restore must discard soft-pin state"; + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + soft_pin_baseline) + << "Restored soft pins must not remain in the active gauge"; + EXPECT_EQ(SoftPinRegistrationCount(restored_service.get()), 0u) + << "Restored soft pins must not remain in the deadline index"; restored_service.reset(); } diff --git a/mooncake-store/tests/ha/snapshot/snapshot_test_utils.h b/mooncake-store/tests/ha/snapshot/snapshot_test_utils.h index d6f0c28ef8..f59c9cd4a4 100644 --- a/mooncake-store/tests/ha/snapshot/snapshot_test_utils.h +++ b/mooncake-store/tests/ha/snapshot/snapshot_test_utils.h @@ -3,6 +3,7 @@ #include #include #include +#include #include #include #include @@ -176,7 +177,8 @@ inline std::vector BuildMetadataPayloadWithClientIdString( uint64_t object_size = kDefaultTestObjectSize, uint64_t put_start_time_ms = kDefaultTestPutStartTimeMs, uint64_t lease_timeout_ms = kDefaultTestLeaseTimeoutMs, - SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy) { + SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy, + std::optional soft_pin_deadline_ms = std::nullopt) { const bool include_data_type = format == SnapshotMetadataFormat::kDataTypeOnly || format == SnapshotMetadataFormat::kDataTypeAndHardPinned || @@ -212,8 +214,8 @@ inline std::vector BuildMetadataPayloadWithClientIdString( shard_packer.pack(put_start_time_ms); shard_packer.pack(object_size); shard_packer.pack(lease_timeout_ms); - shard_packer.pack(false); - shard_packer.pack(uint64_t{0}); + shard_packer.pack(soft_pin_deadline_ms.has_value()); + shard_packer.pack(soft_pin_deadline_ms.value_or(0)); shard_packer.pack(kReplicaCount); if (include_data_type) { shard_packer.pack(static_cast(ObjectDataType::TENSOR)); @@ -238,10 +240,11 @@ inline std::vector BuildMetadataPayload( uint64_t object_size = kDefaultTestObjectSize, uint64_t put_start_time_ms = kDefaultTestPutStartTimeMs, uint64_t lease_timeout_ms = kDefaultTestLeaseTimeoutMs, - SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy) { + SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy, + std::optional soft_pin_deadline_ms = std::nullopt) { return BuildMetadataPayloadWithClientIdString( UuidToString(client_id), object_key, disk_file_path, object_size, - put_start_time_ms, lease_timeout_ms, format); + put_start_time_ms, lease_timeout_ms, format, soft_pin_deadline_ms); } // Builds a metadata payload whose declared replica_count field is set to @@ -361,12 +364,14 @@ inline tl::expected PublishSnapshotPayload( std::string_view object_key = kDefaultTestObjectKey, std::string_view disk_file_path = kDefaultTestDiskFilePath, uint64_t object_size = kDefaultTestObjectSize, - SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy) { + SnapshotMetadataFormat format = SnapshotMetadataFormat::kLegacy, + std::optional soft_pin_deadline_ms = std::nullopt) { return PublishSnapshotPayloadBytes( object_store, catalog_store, descriptor, BuildMetadataPayload(client_id, object_key, disk_file_path, object_size, kDefaultTestPutStartTimeMs, - kDefaultTestLeaseTimeoutMs, format)); + kDefaultTestLeaseTimeoutMs, format, + soft_pin_deadline_ms)); } } // namespace mooncake::test diff --git a/mooncake-store/tests/master_service_test.cpp b/mooncake-store/tests/master_service_test.cpp index 9068dc156f..920d86e1c0 100644 --- a/mooncake-store/tests/master_service_test.cpp +++ b/mooncake-store/tests/master_service_test.cpp @@ -14,6 +14,7 @@ #include #include #include +#include #include #include #include @@ -87,6 +88,77 @@ class MasterServiceTest : public ::testing::Test { static constexpr size_t kDefaultSegmentSize = 1024 * 1024 * 16; static constexpr uint64_t kStrictTenantQuotaBytes = 4 * 1024 * 1024; + std::optional GetSoftPinDeadline( + MasterService& service, const std::string& key, + const std::string& tenant_id = "default") { + const TenantId normalized_tenant = + service.ResolveRequestTenantId(TenantId(tenant_id)); + const size_t shard_idx = + service.getMetadataShardIndex(normalized_tenant, key); + MasterService::MetadataShardAccessorRO shard(&service, shard_idx); + const auto tenant_it = shard->tenants.find(normalized_tenant); + if (tenant_it == shard->tenants.end()) { + return std::nullopt; + } + const auto metadata_it = tenant_it->second.metadata.find(key); + if (metadata_it == tenant_it->second.metadata.end()) { + return std::nullopt; + } + return metadata_it->second.GetCommittedSoftPinTimeout(); + } + + void CleanupExpiredSoftPinsAt( + MasterService& service, + const std::chrono::system_clock::time_point& now) { + service.CleanupExpiredSoftPins(now); + } + + void SetSoftPinDeadlineForTest( + MasterService& service, const std::string& key, + const std::chrono::system_clock::time_point& deadline, + const std::string& tenant_id = "default") { + const TenantId normalized_tenant = + service.ResolveRequestTenantId(TenantId(tenant_id)); + const size_t shard_idx = + service.getMetadataShardIndex(normalized_tenant, key); + MasterService::MetadataShardAccessorRW shard(&service, shard_idx); + auto& metadata = shard->tenants.at(normalized_tenant).metadata.at(key); + { + SpinLocker locker(&metadata.lock); + metadata.soft_pin_timeout = deadline; + } + service.soft_pin_deadline_index_.Upsert( + normalized_tenant.MakeScopedKey(key), shard_idx, deadline); + } + + size_t SoftPinDeadlineHeapSize(MasterService& service) { + return service.soft_pin_deadline_index_.HeapSizeForTest(); + } + + size_t SoftPinRegistrationCount(MasterService& service) { + return service.soft_pin_deadline_index_.RegistrationCountForTest(); + } + + void UpsertSoftPinDeadlineIndexForTest( + MasterService& service, const std::string& key, size_t shard_idx, + const std::chrono::system_clock::time_point& deadline, + const std::string& tenant_id = "default") { + service.soft_pin_deadline_index_.Upsert( + TenantId(tenant_id).MakeScopedKey(key), shard_idx, deadline); + } + + size_t PopExpiredSoftPinDeadlinesForTest( + MasterService& service, + const std::chrono::system_clock::time_point& now) { + return service.soft_pin_deadline_index_.PopExpired(now).size(); + } + + std::chrono::system_clock::time_point ComputeSoftPinDeadlineForTest( + const std::chrono::system_clock::time_point& now, uint64_t ttl_ms) { + return MasterService::ObjectMetadata::ComputeSoftPinDeadline(now, + ttl_ms); + } + std::string WriteTenantPolicyFile( const std::map& tenant_quotas) { TenantQuotaPolicySnapshot snapshot; @@ -668,6 +740,83 @@ TEST_F(MasterServiceTest, PutStartInvalidParams) { EXPECT_EQ(ErrorCode::INVALID_PARAMS, put_result3.error()); } +TEST_F(MasterServiceTest, SoftPinRequestValidation) { + auto service_config = MasterServiceConfig::builder() + .set_default_kv_soft_pin_ttl(50) + .set_max_kv_soft_pin_ttl(100) + .build(); + std::unique_ptr service(new MasterService(service_config)); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + + ReplicateConfig config; + config.soft_pin_ttl_ms = 10; + auto preserve_with_ttl = service->PutStart( + client_id, "preserve_with_ttl", TenantId::Default(), 1024, config); + ASSERT_FALSE(preserve_with_ttl.has_value()); + EXPECT_EQ(preserve_with_ttl.error(), ErrorCode::INVALID_PARAMS); + + config.soft_pin_action = SoftPinAction::DISABLE; + auto disable_with_ttl = service->PutStart( + client_id, "disable_with_ttl", TenantId::Default(), 1024, config); + ASSERT_FALSE(disable_with_ttl.has_value()); + EXPECT_EQ(disable_with_ttl.error(), ErrorCode::INVALID_PARAMS); + + config.soft_pin_action = SoftPinAction::ENABLE; + config.soft_pin_ttl_ms = 101; + auto over_limit = service->PutStart(client_id, "over_limit", + TenantId::Default(), 1024, config); + ASSERT_FALSE(over_limit.has_value()); + EXPECT_EQ(over_limit.error(), ErrorCode::INVALID_PARAMS); + + config.soft_pin_action = static_cast(255); + config.soft_pin_ttl_ms.reset(); + auto invalid_action = service->PutStart(client_id, "invalid_action", + TenantId::Default(), 1024, config); + ASSERT_FALSE(invalid_action.has_value()); + EXPECT_EQ(invalid_action.error(), ErrorCode::INVALID_PARAMS); + + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + config.soft_pin_action = SoftPinAction::ENABLE; + config.soft_pin_ttl_ms = 0; + ASSERT_TRUE( + service + ->PutStart(client_id, "zero_ttl", TenantId::Default(), 1024, config) + .has_value()); + ASSERT_TRUE(service + ->PutEnd(client_id, "zero_ttl", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "zero_ttl").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); +} + +TEST_F(MasterServiceTest, SoftPinMasterConfigRejectsDefaultAboveMaximum) { + auto invalid_config = MasterServiceConfig::builder() + .set_default_kv_soft_pin_ttl(101) + .set_max_kv_soft_pin_ttl(100) + .build(); + EXPECT_THROW(MasterService service(invalid_config), std::invalid_argument); +} + +TEST_F(MasterServiceTest, SoftPinDeadlineCalculationSaturatesAtMaximum) { + using Clock = std::chrono::system_clock; + + const auto normal_now = Clock::time_point(std::chrono::seconds(10)); + EXPECT_EQ(ComputeSoftPinDeadlineForTest(normal_now, 25), + normal_now + std::chrono::milliseconds(25)); + EXPECT_EQ(ComputeSoftPinDeadlineForTest( + normal_now, std::numeric_limits::max()), + Clock::time_point::max()); + + const auto near_max = + Clock::time_point::max() - std::chrono::milliseconds(5); + EXPECT_EQ(ComputeSoftPinDeadlineForTest(near_max, 10), + Clock::time_point::max()); +} + #ifdef USE_NOF TEST_F(MasterServiceTest, PutEndAllCompletesMemoryAndNoFReplicas) { std::unique_ptr service_(new MasterService()); @@ -746,6 +895,73 @@ TEST_F(MasterServiceTest, PutEndMemoryDoesNotCompleteNoFReplica) { ReplicaStatus::COMPLETE); } +TEST_F(MasterServiceTest, PartialRevokePreservesPendingSoftPin) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto mem_context = PrepareSimpleSegment(*service); + NoFSegment nof_segment = + MakeNoFSegment("soft_pin_nof", "soft_pin_nof_endpoint"); + const UUID client_id = generate_uuid(); + ASSERT_TRUE(service->MountNoFSegment(nof_segment, client_id).has_value()); + + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + ReplicateConfig config; + config.replica_num = 1; + config.nof_replica_num = 1; + config.soft_pin_action = SoftPinAction::ENABLE; + ASSERT_TRUE(service + ->PutStart(client_id, "partial_revoke_soft_pin", + TenantId::Default(), 1024, config) + .has_value()); + ASSERT_TRUE(service + ->PutRevoke(client_id, "partial_revoke_soft_pin", + TenantId::Default(), ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + + ASSERT_TRUE(service + ->PutEnd(client_id, "partial_revoke_soft_pin", + TenantId::Default(), ReplicaType::NOF_SSD) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + EXPECT_TRUE( + GetSoftPinDeadline(*service, "partial_revoke_soft_pin").has_value()); +} + +TEST_F(MasterServiceTest, LaterReplicaEndDoesNotRefreshSoftPin) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto mem_context = PrepareSimpleSegment(*service); + NoFSegment nof_segment = + MakeNoFSegment("soft_pin_later_end", "soft_pin_later_end_endpoint"); + const UUID client_id = generate_uuid(); + ASSERT_TRUE(service->MountNoFSegment(nof_segment, client_id).has_value()); + + ReplicateConfig config; + config.replica_num = 1; + config.nof_replica_num = 1; + config.soft_pin_action = SoftPinAction::ENABLE; + ASSERT_TRUE(service + ->PutStart(client_id, "later_end_soft_pin", + TenantId::Default(), 1024, config) + .has_value()); + ASSERT_TRUE(service + ->PutEnd(client_id, "later_end_soft_pin", + TenantId::Default(), ReplicaType::MEMORY) + .has_value()); + const auto first_deadline = + GetSoftPinDeadline(*service, "later_end_soft_pin"); + ASSERT_TRUE(first_deadline.has_value()); + + ASSERT_TRUE(service + ->PutEnd(client_id, "later_end_soft_pin", + TenantId::Default(), ReplicaType::NOF_SSD) + .has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "later_end_soft_pin"), + first_deadline); +} + TEST_F(MasterServiceTest, PutStartOnePlusOneAllowsSingleAllocatedReplica) { std::unique_ptr service_(new MasterService()); [[maybe_unused]] const auto mem_context = PrepareSimpleSegment(*service_); @@ -4429,7 +4645,7 @@ TEST_F(MasterServiceTest, RemoveSoftPinObject) { uint64_t slice_length = 1024; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; // Verify soft pin does not block remove ASSERT_TRUE(service_ @@ -4440,7 +4656,9 @@ TEST_F(MasterServiceTest, RemoveSoftPinObject) { service_ ->PutEnd(client_id, key, TenantId::Default(), ReplicaType::MEMORY) .has_value()); + EXPECT_EQ(SoftPinRegistrationCount(*service_), 1u); EXPECT_TRUE(service_->Remove(key, TenantId::Default()).has_value()); + EXPECT_EQ(SoftPinRegistrationCount(*service_), 0u); // Verify soft pin does not block RemoveAll ASSERT_TRUE(service_ @@ -4451,7 +4669,322 @@ TEST_F(MasterServiceTest, RemoveSoftPinObject) { service_ ->PutEnd(client_id, key, TenantId::Default(), ReplicaType::MEMORY) .has_value()); + EXPECT_EQ(SoftPinRegistrationCount(*service_), 1u); EXPECT_EQ(1, service_->RemoveAll()); + EXPECT_EQ(SoftPinRegistrationCount(*service_), 0u); +} + +TEST_F(MasterServiceTest, SoftPinActionsCommitOnFirstReadableUpsert) { + auto service_config = MasterServiceConfig::builder() + .set_default_kv_soft_pin_ttl(10000) + .build(); + std::unique_ptr service(new MasterService(service_config)); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + + ReplicateConfig enable; + enable.soft_pin_action = SoftPinAction::ENABLE; + enable.soft_pin_ttl_ms = 5000; + ASSERT_TRUE(service + ->PutStart(client_id, "action_key", TenantId::Default(), + 1024, enable) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + const auto before_first_completion = std::chrono::system_clock::now(); + ASSERT_TRUE(service + ->PutEnd(client_id, "action_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + const auto initial_deadline = GetSoftPinDeadline(*service, "action_key"); + ASSERT_TRUE(initial_deadline.has_value()); + EXPECT_GT(*initial_deadline, + before_first_completion + std::chrono::seconds(4)); + EXPECT_LT(*initial_deadline, + before_first_completion + std::chrono::seconds(9)); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + + ReplicateConfig preserve; + ASSERT_TRUE(service + ->UpsertStart(client_id, "action_key", TenantId::Default(), + 1024, preserve) + .has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "action_key"), initial_deadline); + ASSERT_TRUE(service + ->UpsertEnd(client_id, "action_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "action_key"), initial_deadline); + + ReplicateConfig disable; + disable.soft_pin_action = SoftPinAction::DISABLE; + ASSERT_TRUE(service + ->UpsertStart(client_id, "action_key", TenantId::Default(), + 2048, disable) + .has_value()); + EXPECT_TRUE(GetSoftPinDeadline(*service, "action_key").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + ASSERT_TRUE(service + ->UpsertEnd(client_id, "action_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "action_key").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + + ReplicateConfig enable_again; + enable_again.soft_pin_action = SoftPinAction::ENABLE; + enable_again.soft_pin_ttl_ms = 3000; + ASSERT_TRUE(service + ->UpsertStart(client_id, "action_key", TenantId::Default(), + 2048, enable_again) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "action_key").has_value()); + const auto before_enable_again = std::chrono::system_clock::now(); + ASSERT_TRUE(service + ->UpsertEnd(client_id, "action_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + const auto enabled_again_deadline = + GetSoftPinDeadline(*service, "action_key"); + ASSERT_TRUE(enabled_again_deadline.has_value()); + EXPECT_GT(*enabled_again_deadline, + before_enable_again + std::chrono::seconds(2)); + EXPECT_LT(*enabled_again_deadline, + before_enable_again + std::chrono::seconds(5)); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); +} + +TEST_F(MasterServiceTest, SoftPinDeadlineIndexExpiresOnlyDueEntries) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + + ReplicateConfig config; + config.soft_pin_action = SoftPinAction::ENABLE; + PutCompletedObject(*service, client_id, "deadline_key", config); + + ReplicateConfig grouped_config = config; + grouped_config.group_ids = std::vector{ + FindGroupIdOnDifferentShard("grouped_deadline_key")}; + PutCompletedObject(*service, client_id, "grouped_deadline_key", + grouped_config); + + const auto first_deadline = + std::chrono::system_clock::now() + std::chrono::hours(1); + const auto second_deadline = first_deadline + std::chrono::seconds(1); + SetSoftPinDeadlineForTest(*service, "deadline_key", first_deadline); + SetSoftPinDeadlineForTest(*service, "grouped_deadline_key", + second_deadline); + + EXPECT_EQ(SoftPinRegistrationCount(*service), 2u); + CleanupExpiredSoftPinsAt(*service, first_deadline); + EXPECT_FALSE(GetSoftPinDeadline(*service, "deadline_key").has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "grouped_deadline_key"), + second_deadline); + EXPECT_EQ(SoftPinRegistrationCount(*service), 1u); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + + CleanupExpiredSoftPinsAt(*service, second_deadline); + EXPECT_FALSE( + GetSoftPinDeadline(*service, "grouped_deadline_key").has_value()); + EXPECT_EQ(SoftPinRegistrationCount(*service), 0u); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); +} + +TEST_F(MasterServiceTest, SoftPinTtlUpdateInvalidatesOldHeapEntry) { + auto service_config = MasterServiceConfig::builder() + .set_default_kv_soft_pin_ttl(5000) + .build(); + std::unique_ptr service(new MasterService(service_config)); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + + ReplicateConfig enable; + enable.soft_pin_action = SoftPinAction::ENABLE; + PutCompletedObject(*service, client_id, "ttl_update_key", enable); + const auto first_deadline = GetSoftPinDeadline(*service, "ttl_update_key"); + ASSERT_TRUE(first_deadline.has_value()); + + enable.soft_pin_ttl_ms = 20000; + ASSERT_TRUE(service + ->UpsertStart(client_id, "ttl_update_key", + TenantId::Default(), 1024, enable) + .has_value()); + ASSERT_TRUE(service + ->UpsertEnd(client_id, "ttl_update_key", + TenantId::Default(), ReplicaType::MEMORY) + .has_value()); + const auto updated_deadline = + GetSoftPinDeadline(*service, "ttl_update_key"); + ASSERT_TRUE(updated_deadline.has_value()); + EXPECT_GT(*updated_deadline, *first_deadline); + EXPECT_EQ(SoftPinRegistrationCount(*service), 1u); + EXPECT_GE(SoftPinDeadlineHeapSize(*service), 2u); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + + CleanupExpiredSoftPinsAt(*service, *first_deadline); + EXPECT_EQ(GetSoftPinDeadline(*service, "ttl_update_key"), updated_deadline); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); + + CleanupExpiredSoftPinsAt(*service, *updated_deadline); + EXPECT_FALSE(GetSoftPinDeadline(*service, "ttl_update_key").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); +} + +TEST_F(MasterServiceTest, + SizeChangingUpsertIndexesInheritedDeadlineBeforeCompletion) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + + ReplicateConfig enable; + enable.soft_pin_action = SoftPinAction::ENABLE; + PutCompletedObject(*service, client_id, "resize_pending", enable); + const auto inherited_deadline = + std::chrono::system_clock::now() + std::chrono::hours(1); + SetSoftPinDeadlineForTest(*service, "resize_pending", inherited_deadline); + + ReplicateConfig preserve; + ASSERT_TRUE(service + ->UpsertStart(client_id, "resize_pending", + TenantId::Default(), 2048, preserve) + .has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "resize_pending"), + inherited_deadline); + EXPECT_EQ(SoftPinRegistrationCount(*service), 1u); + + CleanupExpiredSoftPinsAt(*service, inherited_deadline); + EXPECT_FALSE(GetSoftPinDeadline(*service, "resize_pending").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + + ASSERT_TRUE(service + ->UpsertEnd(client_id, "resize_pending", + TenantId::Default(), ReplicaType::MEMORY) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "resize_pending").has_value()); + EXPECT_EQ(SoftPinRegistrationCount(*service), 0u); +} + +TEST_F(MasterServiceTest, SoftPinDeadlineHeapCompactsRepeatedUpdates) { + MasterService service; + const auto base = std::chrono::system_clock::now(); + constexpr size_t kUpdates = 5000; + for (size_t i = 0; i < kUpdates; ++i) { + UpsertSoftPinDeadlineIndexForTest( + service, "compaction_key", 0, + base + std::chrono::milliseconds(i + 1)); + } + + EXPECT_EQ(SoftPinRegistrationCount(service), 1u); + EXPECT_LE(SoftPinDeadlineHeapSize(service), 4096u); + EXPECT_EQ(PopExpiredSoftPinDeadlinesForTest( + service, base + std::chrono::milliseconds(kUpdates - 1)), + 0u); + EXPECT_EQ(SoftPinRegistrationCount(service), 1u); + EXPECT_EQ(PopExpiredSoftPinDeadlinesForTest( + service, base + std::chrono::milliseconds(kUpdates)), + 1u); +} + +TEST_F(MasterServiceTest, + ExpiredSoftPinIsNotCarriedAcrossUpsertMetadataReplacement) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_a = generate_uuid(); + const UUID client_b = generate_uuid(); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + + ReplicateConfig enable; + enable.soft_pin_action = SoftPinAction::ENABLE; + enable.soft_pin_ttl_ms = 10000; + ReplicateConfig preserve; + + ASSERT_TRUE( + service + ->PutStart(client_a, "preempted", TenantId::Default(), 1024, enable) + .has_value()); + ASSERT_TRUE(service + ->PutEnd(client_a, "preempted", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + SetSoftPinDeadlineForTest( + *service, "preempted", + std::chrono::system_clock::now() - std::chrono::seconds(1)); + ASSERT_TRUE(service + ->UpsertStart(client_a, "preempted", TenantId::Default(), + 1024, preserve) + .has_value()); + ASSERT_TRUE(service + ->UpsertStart(client_b, "preempted", TenantId::Default(), + 1024, preserve) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "preempted").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + + ASSERT_TRUE( + service + ->PutStart(client_a, "resized", TenantId::Default(), 1024, enable) + .has_value()); + ASSERT_TRUE(service + ->PutEnd(client_a, "resized", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + SetSoftPinDeadlineForTest( + *service, "resized", + std::chrono::system_clock::now() - std::chrono::seconds(1)); + ASSERT_TRUE(service + ->UpsertStart(client_a, "resized", TenantId::Default(), + 2048, preserve) + .has_value()); + EXPECT_FALSE(GetSoftPinDeadline(*service, "resized").has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); +} + +TEST_F(MasterServiceTest, RepeatedPutEndDoesNotRefreshSoftPin) { + std::unique_ptr service(new MasterService()); + [[maybe_unused]] const auto context = PrepareSimpleSegment(*service); + const UUID client_id = generate_uuid(); + + ReplicateConfig config; + config.soft_pin_action = SoftPinAction::ENABLE; + config.soft_pin_ttl_ms = 5000; + ASSERT_TRUE(service + ->PutStart(client_id, "repeat_end", TenantId::Default(), + 1024, config) + .has_value()); + ASSERT_TRUE(service + ->PutEnd(client_id, "repeat_end", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + const auto first_deadline = GetSoftPinDeadline(*service, "repeat_end"); + ASSERT_TRUE(first_deadline.has_value()); + + ASSERT_TRUE(service + ->PutEnd(client_id, "repeat_end", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(GetSoftPinDeadline(*service, "repeat_end"), first_deadline); } TEST_F(MasterServiceTest, SoftPinObjectsNotEvictedBeforeOtherObjects) { @@ -4485,7 +5018,7 @@ TEST_F(MasterServiceTest, SoftPinObjectsNotEvictedBeforeOtherObjects) { uint64_t slice_length = value_size; ReplicateConfig soft_pin_config; soft_pin_config.replica_num = 1; - soft_pin_config.with_soft_pin = true; + soft_pin_config.soft_pin_action = SoftPinAction::ENABLE; ASSERT_TRUE(service_ ->PutStart(client_id, pin_key, TenantId::Default(), @@ -4562,7 +5095,7 @@ TEST_F(MasterServiceTest, SoftPinObjectsCanBeEvicted) { uint64_t slice_length = value_size; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; if (service_ ->PutStart(client_id, key, TenantId::Default(), slice_length, config) @@ -4582,98 +5115,51 @@ TEST_F(MasterServiceTest, SoftPinObjectsCanBeEvicted) { service_->RemoveAll(); } -TEST_F(MasterServiceTest, SoftPinExtendedOnGet) { +TEST_F(MasterServiceTest, SoftPinExpiresAndGetDoesNotReactivate) { const uint64_t kv_lease_ttl = 200; - // The soft pin ttl shall not be too large, otherwise the test will take too - // long - const uint64_t kv_soft_pin_ttl = 1000; - static_assert( - kv_soft_pin_ttl > kv_lease_ttl, - "kv_soft_pin_ttl must be larger than kv_lease_ttl in this test"); - const double eviction_ratio = 0.5; - const bool allow_evict_soft_pinned_objects = true; + const uint64_t kv_soft_pin_ttl = 20; auto service_config = MasterServiceConfig::builder() .set_default_kv_lease_ttl(kv_lease_ttl) .set_default_kv_soft_pin_ttl(kv_soft_pin_ttl) - .set_allow_evict_soft_pinned_objects( - allow_evict_soft_pinned_objects) - .set_eviction_ratio(eviction_ratio) .build(); std::unique_ptr service_(new MasterService(service_config)); const UUID client_id = generate_uuid(); - // Mount segment and put an object constexpr size_t buffer = 0x300000000; constexpr size_t segment_size = 1024 * 1024 * 16; - constexpr size_t value_size = 1024 * 1024; + constexpr size_t value_size = 1024; [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_, "test_segment", buffer, segment_size); - // The eviction has random factors, so test 3 times - for (int test_i = 0; test_i < 3; test_i++) { - // Put pin_key first - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - uint64_t slice_length = value_size; - ReplicateConfig soft_pin_config; - soft_pin_config.replica_num = 1; - soft_pin_config.with_soft_pin = true; - - ASSERT_TRUE(service_->PutStart(client_id, pin_key, - TenantId::Default(), slice_length, - soft_pin_config)); - ASSERT_TRUE(service_ - ->PutEnd(client_id, pin_key, TenantId::Default(), - ReplicaType::MEMORY) - .has_value()); - } - - // Wait for the soft pin to expire - std::this_thread::sleep_for(std::chrono::milliseconds(kv_soft_pin_ttl)); - - // Get the pin_key to extend the soft pin - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - ASSERT_TRUE(service_->GetReplicaList(pin_key, TenantId::Default()) - .has_value()); - } - - // Fill the segment to trigger eviction - int failed_puts = 0; - for (int i = 0; i < 16; i++) { - std::string key = "key" + std::to_string(i); - uint64_t slice_length = value_size; - ReplicateConfig config; - config.replica_num = 1; - if (service_ - ->PutStart(client_id, key, TenantId::Default(), - slice_length, config) - .has_value()) { - ASSERT_TRUE(service_ - ->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY) - .has_value()); - } else { - failed_puts++; - } - } - ASSERT_GT(failed_puts, 0); - - // wait for eviction - std::this_thread::sleep_for(std::chrono::milliseconds(kv_lease_ttl)); + const int64_t baseline = + MasterMetricManager::instance().get_soft_pin_key_count(); + ReplicateConfig config; + config.soft_pin_action = SoftPinAction::ENABLE; + ASSERT_TRUE(service_ + ->PutStart(client_id, "pin_key", TenantId::Default(), + value_size, config) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + ASSERT_TRUE(service_ + ->PutEnd(client_id, "pin_key", TenantId::Default(), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline + 1); - // pin_key should still be accessible - for (int i = 0; i < 2; i++) { - std::string pin_key = "pin_key" + std::to_string(i); - ASSERT_TRUE(service_->GetReplicaList(pin_key, TenantId::Default()) - .has_value()); - } + const auto deadline = GetSoftPinDeadline(*service_, "pin_key"); + ASSERT_TRUE(deadline.has_value()); + CleanupExpiredSoftPinsAt(*service_, *deadline); + ASSERT_TRUE( + service_->GetReplicaList("pin_key", TenantId::Default()).has_value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); - // wait for the lease to expire - std::this_thread::sleep_for(std::chrono::milliseconds(kv_lease_ttl)); - // remove all objects before the next turn - service_->RemoveAll(); - } + ASSERT_TRUE(service_->ExistKey("pin_key", TenantId::Default()).value()); + EXPECT_EQ(MasterMetricManager::instance().get_soft_pin_key_count(), + baseline); + service_->RemoveAll(); } TEST_F(MasterServiceTest, SoftPinObjectsNotAllowEvict) { @@ -4706,7 +5192,7 @@ TEST_F(MasterServiceTest, SoftPinObjectsNotAllowEvict) { uint64_t slice_length = value_size; ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; if (service_ ->PutStart(client_id, key, TenantId::Default(), slice_length, config) @@ -5668,7 +6154,7 @@ TEST_F(MasterServiceTest, BatchReplicaClearSpecificSegment) { ASSERT_TRUE(put_end_result.has_value()); // 4. Wait for lease to expire and verify it's actually expired - // PutEnd calls GrantLease(0, ...) which sets lease_timeout to now. + // PutEnd grants a zero-duration read lease, setting lease_timeout to now. // Due to clock precision and timing, we need to ensure the lease is // actually expired before calling BatchReplicaClear. // Use a small delay and then poll to ensure lease is expired. @@ -7086,7 +7572,7 @@ TEST_F(MasterServiceTest, HardPinWithSoftPinEvictionOrder) { { ReplicateConfig config; config.replica_num = 1; - config.with_soft_pin = true; + config.soft_pin_action = SoftPinAction::ENABLE; ASSERT_TRUE(service_ ->PutStart(client_id, "soft_pinned", TenantId::Default(), value_size, config) diff --git a/mooncake-wheel/tests/test_distributed_object_store.py b/mooncake-wheel/tests/test_distributed_object_store.py index 880a3b532c..bbb967ec6b 100644 --- a/mooncake-wheel/tests/test_distributed_object_store.py +++ b/mooncake-wheel/tests/test_distributed_object_store.py @@ -2,7 +2,8 @@ import os import time import threading -from mooncake.store import MooncakeDistributedStore +import random +from mooncake.store import MooncakeDistributedStore, SoftPinAction # The lease time of the kv object, should be set equal to # the master's value. @@ -53,6 +54,27 @@ def get_config_dict(global_segment_size, local_buffer_size): } +class TestReplicateConfig(unittest.TestCase): + """Test ReplicateConfig bindings without a running store service.""" + + def test_soft_pin_ttl_optional_conversion(self): + from mooncake.store import ReplicateConfig + + config = ReplicateConfig() + self.assertIsNone(config.soft_pin_ttl_ms) + + config.soft_pin_ttl_ms = 1000 + self.assertEqual(config.soft_pin_ttl_ms, 1000) + + config.soft_pin_ttl_ms = None + self.assertIsNone(config.soft_pin_ttl_ms) + + with self.assertRaises(TypeError): + config.soft_pin_ttl_ms = -1 + with self.assertRaises(TypeError): + config.soft_pin_ttl_ms = 1 << 64 + + class TestConfigDictSetup(unittest.TestCase): """Test configuration-dictionary setup through the Python store wrapper.""" @@ -156,6 +178,33 @@ def test_basic_put_get_exist_operations(self): time.sleep(default_kv_lease_ttl / 1000) self.assertEqual(self.store.remove(key), 0) + def test_soft_pin_config_forwarding(self): + """Test soft-pin action and TTL forwarding through Store operations.""" + from mooncake.store import ReplicateConfig + + key = f"test_soft_pin_config_forwarding_{os.getpid()}" + test_data = b"soft-pin forwarding" + + enable_config = ReplicateConfig() + enable_config.soft_pin_action = SoftPinAction.ENABLE + enable_config.soft_pin_ttl_ms = 1000 + self.assertEqual(self.store.put(key, test_data, enable_config), 0) + + preserve_with_ttl = ReplicateConfig() + preserve_with_ttl.soft_pin_ttl_ms = 1000 + self.assertEqual(self.store.upsert(key, test_data, preserve_with_ttl), -600) + + disable_with_ttl = ReplicateConfig() + disable_with_ttl.soft_pin_action = SoftPinAction.DISABLE + disable_with_ttl.soft_pin_ttl_ms = 1000 + self.assertEqual(self.store.upsert(key, test_data, disable_with_ttl), -600) + + self.assertEqual(self.store.upsert(key, test_data), 0) + + disable_config = ReplicateConfig() + disable_config.soft_pin_action = SoftPinAction.DISABLE + self.assertEqual(self.store.upsert(key, test_data, disable_config), 0) + def test_batch_is_exist_operations(self): """Test batch is_exist operations through the Python interface.""" batch_size = 20 @@ -662,16 +711,16 @@ def test_replicate_config_creation_and_properties(self): # Test default constructor config = ReplicateConfig() self.assertEqual(config.replica_num, 1) - self.assertEqual(config.with_soft_pin, False) + self.assertEqual(config.soft_pin_action, SoftPinAction.PRESERVE) self.assertEqual(config.preferred_segment, "") # Test property assignment config.replica_num = 3 - config.with_soft_pin = True + config.soft_pin_action = SoftPinAction.ENABLE config.preferred_segment = "node1:12345" self.assertEqual(config.replica_num, 3) - self.assertEqual(config.with_soft_pin, True) + self.assertEqual(config.soft_pin_action, SoftPinAction.ENABLE) self.assertEqual(config.preferred_segment, "node1:12345") # Test string representation diff --git a/mooncake-wheel/tests/test_distributed_object_store_cxl.py b/mooncake-wheel/tests/test_distributed_object_store_cxl.py index a08382f5dc..e41e6f5c49 100644 --- a/mooncake-wheel/tests/test_distributed_object_store_cxl.py +++ b/mooncake-wheel/tests/test_distributed_object_store_cxl.py @@ -3,7 +3,7 @@ import time import threading import tempfile -from mooncake.store import MooncakeDistributedStore +from mooncake.store import MooncakeDistributedStore, SoftPinAction CXL_SIM_FILE = os.path.join(tempfile.gettempdir(), "tmp_dax_sim") CXL_SIM_SIZE = 8 * 1024 * 1024 * 1024 @@ -588,16 +588,16 @@ def test_replicate_config_creation_and_properties(self): # Test default constructor config = ReplicateConfig() self.assertEqual(config.replica_num, 1) - self.assertEqual(config.with_soft_pin, False) + self.assertEqual(config.soft_pin_action, SoftPinAction.PRESERVE) self.assertEqual(config.preferred_segment, "") # Test property assignment config.replica_num = 3 - config.with_soft_pin = True + config.soft_pin_action = SoftPinAction.ENABLE config.preferred_segment = "node1:12345" self.assertEqual(config.replica_num, 3) - self.assertEqual(config.with_soft_pin, True) + self.assertEqual(config.soft_pin_action, SoftPinAction.ENABLE) self.assertEqual(config.preferred_segment, "node1:12345") # Test string representation diff --git a/mooncake-wheel/tests/test_dummy_client.py b/mooncake-wheel/tests/test_dummy_client.py index c35c9d9492..bbc0adddaf 100644 --- a/mooncake-wheel/tests/test_dummy_client.py +++ b/mooncake-wheel/tests/test_dummy_client.py @@ -8,7 +8,7 @@ except ImportError: torch = None -from mooncake.store import MooncakeDistributedStore +from mooncake.store import MooncakeDistributedStore, SoftPinAction # The lease time of the kv object, should be set equal to # the master's value. @@ -626,16 +626,16 @@ def test_replicate_config_creation_and_properties(self): # Test default constructor config = ReplicateConfig() self.assertEqual(config.replica_num, 1) - self.assertEqual(config.with_soft_pin, False) + self.assertEqual(config.soft_pin_action, SoftPinAction.PRESERVE) self.assertEqual(config.preferred_segment, "") # Test property assignment config.replica_num = 3 - config.with_soft_pin = True + config.soft_pin_action = SoftPinAction.ENABLE config.preferred_segment = "node1:12345" self.assertEqual(config.replica_num, 3) - self.assertEqual(config.with_soft_pin, True) + self.assertEqual(config.soft_pin_action, SoftPinAction.ENABLE) self.assertEqual(config.preferred_segment, "node1:12345") # Test string representation diff --git a/mooncake-wheel/tests/test_replicated_distributed_object_store.py b/mooncake-wheel/tests/test_replicated_distributed_object_store.py index 0ad91d3e6a..d4922a0d97 100644 --- a/mooncake-wheel/tests/test_replicated_distributed_object_store.py +++ b/mooncake-wheel/tests/test_replicated_distributed_object_store.py @@ -1,7 +1,7 @@ import unittest import os import time -from mooncake.store import MooncakeDistributedStore, ReplicateConfig +from mooncake.store import MooncakeDistributedStore, ReplicateConfig, SoftPinAction # The lease time of the kv object, should be set equal to # the master's value. @@ -173,7 +173,7 @@ def test_put_from_with_config_parameter(self): # Test with custom config config = ReplicateConfig() config.replica_num = self.max_replicate_num - config.with_soft_pin = False + config.soft_pin_action = SoftPinAction.PRESERVE config.with_hard_pin = False key2 = "test_put_from_config_key2" @@ -244,7 +244,7 @@ def test_batch_put_from_with_config_parameter(self): # Test with custom config config = ReplicateConfig() config.replica_num = self.max_replicate_num - config.with_soft_pin = False + config.soft_pin_action = SoftPinAction.PRESERVE config.with_hard_pin = False keys2 = ["test_batch_put_from_config_key4", "test_batch_put_from_config_key5", "test_batch_put_from_config_key6"] From 98547274bfafa8426cb55d37f6803e3975c17083 Mon Sep 17 00:00:00 2001 From: Ola <114643959+030611@users.noreply.github.com> Date: Tue, 11 Aug 2026 12:04:51 +0800 Subject: [PATCH 027/483] [Bugfix] use BufferPool for zcopy benchmark (#3342) * [Bugfix] use BufferPool for zcopy benchmark Fixes kvcache-ai/Mooncake#3329. Use the native BufferPool/BufferLease path instead of the NoF-only hugepage helpers and add CPU-only regression coverage. * Normalize benchmark files to LF line endings Preserve the repository line-ending convention so the PR diff only contains the intended BufferPool fix and regression tests. * [Bugfix] resolve BufferPool only for zcopy --- mooncake-store/benchmarks/store_kv_bench.py | 83 +++++---------- .../benchmarks/test_store_kv_bench.py | 100 +++++++++++++++++- 2 files changed, 125 insertions(+), 58 deletions(-) diff --git a/mooncake-store/benchmarks/store_kv_bench.py b/mooncake-store/benchmarks/store_kv_bench.py index f09151cfb8..bc637d287f 100644 --- a/mooncake-store/benchmarks/store_kv_bench.py +++ b/mooncake-store/benchmarks/store_kv_bench.py @@ -24,17 +24,6 @@ MooncakeDistributedStore = _store_module.MooncakeDistributedStore ReplicateConfig = _store_module.ReplicateConfig -try: - get_alloc_func_addr = _store_module.get_alloc_func_addr - get_free_func_addr = _store_module.get_free_func_addr -except AttributeError: - - def get_alloc_func_addr(): - return None - - def get_free_func_addr(): - return None - LOG = logging.getLogger("store_kv_bench") @@ -592,61 +581,41 @@ def _get_lengths_zcopy(self, keys: List[str], slot_count: int) -> List[int]: class ZcopyBufferPool: - def __init__(self, store_obj, value_size: int, slots: int): + def __init__(self, store_obj, value_size: int, slots: int, local_buffer_size: int): self.store = store_obj self.value_size = value_size self.slots = slots self.total_size = self.value_size * self.slots - self._alloc_fn = None - self._free_fn = None - self._registered = False + self._pool = None + self._lease = None self.base_ptr = 0 - alloc_addr = get_alloc_func_addr() - free_addr = get_free_func_addr() - if alloc_addr is None or free_addr is None: - raise RuntimeError( - "store module does not expose hugepage alloc/free helpers" + if self.total_size > local_buffer_size: + raise ValueError( + "zcopy pool requires " + f"{self.total_size} bytes, but --local-buffer-size is " + f"{local_buffer_size}; increase --local-buffer-size" ) - self._alloc_fn = ctypes.CFUNCTYPE(ctypes.c_void_p, ctypes.c_size_t)( - get_alloc_func_addr() + self._pool = _store_module.BufferPool( + self.store, max_bytes=local_buffer_size, max_regions=1 ) - self._free_fn = ctypes.CFUNCTYPE(None, ctypes.c_void_p)(get_free_func_addr()) - - raw_ptr = self._alloc_fn(self.total_size) - self.base_ptr = ctypes.cast(raw_ptr, ctypes.c_void_p).value or 0 - if self.base_ptr == 0: - raise RuntimeError( - f"direct hugepage alloc failed for zcopy pool: size={self.total_size}" - ) - ret = self.store.register_buffer(self.base_ptr, self.total_size) - if ret != 0: - failed_ptr = self.base_ptr - self._free_fn(ctypes.c_void_p(self.base_ptr)) - self.base_ptr = 0 - raise RuntimeError( - f"register_buffer failed for direct zcopy pool ptr={failed_ptr}: {ret}" - ) - self._registered = True - self._buffer = (ctypes.c_ubyte * self.total_size).from_address(self.base_ptr) + try: + self._lease = self._pool.acquire(self.total_size, block=False) + except Exception: + self._pool.close() + self._pool = None + raise + self.base_ptr = self._lease.ptr def close(self) -> None: - self._buffer = None - if self.base_ptr: - if self._registered: - try: - self.store.unregister_buffer(self.base_ptr) - except Exception: - LOG.debug( - "unregister_buffer failed for direct zcopy pool ptr=%s", - self.base_ptr, - exc_info=True, - ) - self._registered = False - if self._free_fn is not None: - self._free_fn(ctypes.c_void_p(self.base_ptr)) - self.base_ptr = 0 + self.base_ptr = 0 + if self._lease is not None: + self._lease.release() + self._lease = None + if self._pool is not None: + self._pool.close() + self._pool = None def slot_ptr(self, slot: int) -> int: if slot < 0 or slot >= self.slots: @@ -714,7 +683,9 @@ def __init__(self, args: argparse.Namespace, lane_count: int): self.zcopy_pool: Optional[ZcopyBufferPool] = None if args.io_api == "zcopy": slots = max(1, args.batch_size) * lane_count - self.zcopy_pool = ZcopyBufferPool(self.store, args.value_size, slots) + self.zcopy_pool = ZcopyBufferPool( + self.store, args.value_size, slots, args.local_buffer_size + ) def make_session( self, diff --git a/mooncake-store/benchmarks/test_store_kv_bench.py b/mooncake-store/benchmarks/test_store_kv_bench.py index 2ede5f94f4..aedfaee1bf 100644 --- a/mooncake-store/benchmarks/test_store_kv_bench.py +++ b/mooncake-store/benchmarks/test_store_kv_bench.py @@ -1,3 +1,4 @@ +import ctypes import importlib.util import json import os @@ -8,11 +9,42 @@ from types import SimpleNamespace +class FakeBufferLease: + def __init__(self, size): + self._buffer = (ctypes.c_ubyte * size)() + self.ptr = ctypes.addressof(self._buffer) + self.released = False + + def release(self): + self.released = True + + +class FakeBufferPool: + instances = [] + acquire_error = None + + def __init__(self, store, **kwargs): + self.store = store + self.kwargs = kwargs + self.lease = None + self.closed = False + self.instances.append(self) + + def acquire(self, size, block=True): + if self.acquire_error is not None: + raise self.acquire_error + self.lease = FakeBufferLease(size) + self.acquire_args = (size, block) + return self.lease + + def close(self): + self.closed = True + + store_module = types.ModuleType("mooncake.store") store_module.MooncakeDistributedStore = object store_module.ReplicateConfig = type("ReplicateConfig", (), {}) -store_module.get_alloc_func_addr = lambda: None -store_module.get_free_func_addr = lambda: None +store_module.BufferPool = FakeBufferPool sys.modules.setdefault("mooncake", types.ModuleType("mooncake")) sys.modules["mooncake.store"] = store_module @@ -23,6 +55,20 @@ SPEC.loader.exec_module(bench) +class StoreModuleCompatibilityTest(unittest.TestCase): + def test_import_does_not_require_buffer_pool(self): + del store_module.BufferPool + module_name = "store_kv_bench_without_buffer_pool" + spec = importlib.util.spec_from_file_location(module_name, MODULE_PATH) + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + try: + spec.loader.exec_module(module) + finally: + store_module.BufferPool = FakeBufferPool + sys.modules.pop(module_name, None) + + class MetadataWorkloadModelTest(unittest.TestCase): def test_percentages_and_lane_choices_are_deterministic(self): weights = {"put": 40, "get": 30, "exist": 20, "remove": 10} @@ -170,5 +216,55 @@ def remove(self, key): self.assertEqual(summary["journal_records"], 5) +class ZcopyBufferPoolTest(unittest.TestCase): + def setUp(self): + FakeBufferPool.instances.clear() + FakeBufferPool.acquire_error = None + + def test_uses_native_buffer_pool_without_nof_allocators(self): + store = object() + pool = bench.ZcopyBufferPool( + store, value_size=16, slots=2, local_buffer_size=64 + ) + native_pool = FakeBufferPool.instances[0] + + self.assertIs(native_pool.store, store) + self.assertEqual(native_pool.kwargs, {"max_bytes": 64, "max_regions": 1}) + self.assertEqual(native_pool.acquire_args, (32, False)) + self.assertEqual(pool.slot_ptr(1), pool.base_ptr + 16) + + view = bench.ZcopyBufferView(pool, slot_offset=0, slots=2) + self.assertEqual(view.fill_write_buffers([b"abc"]), [pool.base_ptr]) + self.assertEqual(view.read_bytes(0, 3), b"abc") + + lease = native_pool.lease + pool.close() + self.assertTrue(lease.released) + self.assertTrue(native_pool.closed) + self.assertEqual(pool.base_ptr, 0) + + pool.close() + self.assertTrue(lease.released) + self.assertTrue(native_pool.closed) + + def test_rejects_pool_larger_than_store_local_buffer(self): + with self.assertRaisesRegex(ValueError, "increase --local-buffer-size"): + bench.ZcopyBufferPool( + object(), value_size=32, slots=3, local_buffer_size=64 + ) + self.assertEqual(FakeBufferPool.instances, []) + + def test_closes_native_pool_when_acquire_fails(self): + FakeBufferPool.acquire_error = RuntimeError("buffer pool is exhausted") + + with self.assertRaisesRegex(RuntimeError, "buffer pool is exhausted"): + bench.ZcopyBufferPool( + object(), value_size=16, slots=2, local_buffer_size=64 + ) + + self.assertEqual(len(FakeBufferPool.instances), 1) + self.assertTrue(FakeBufferPool.instances[0].closed) + + if __name__ == "__main__": unittest.main() From 3a6c786caa4a2c8aae94a5325a17a281cf142087 Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Tue, 11 Aug 2026 14:07:39 +0800 Subject: [PATCH 028/483] [Store] Avoid staging for same-process GPU reads (#3197) --- mooncake-store/include/client_service.h | 6 ++ mooncake-store/src/real_client.cpp | 113 +++++++++++--------- mooncake-store/tests/transfer_task_test.cpp | 63 +++++++++++ 3 files changed, 134 insertions(+), 48 deletions(-) diff --git a/mooncake-store/include/client_service.h b/mooncake-store/include/client_service.h index 78e324e8be..094f8d8cc6 100644 --- a/mooncake-store/include/client_service.h +++ b/mooncake-store/include/client_service.h @@ -647,6 +647,12 @@ class Client { return endpoints; } + bool CanUseLocalMemcpy(const Replica::Descriptor& replica) const { + if (!replica.is_memory_replica()) return false; + return CanUseLocalMemcpy(replica.get_memory_descriptor() + .buffer_descriptor.transport_endpoint_); + } + /** * @brief Check if local hot cache is enabled * @return true if hot cache is enabled, false otherwise diff --git a/mooncake-store/src/real_client.cpp b/mooncake-store/src/real_client.cpp index a774762a64..df68476af8 100644 --- a/mooncake-store/src/real_client.cpp +++ b/mooncake-store/src/real_client.cpp @@ -3390,7 +3390,9 @@ tl::expected RealClient::execute_ranged_read( auto runtime_accelerator = device::GetAcceleratorRegistry().RuntimeAccelerators(); void *dst = static_cast(buffer) + dst_offset; - if (runtime_accelerator.FindDeviceForPointer(dst)) { + if (runtime_accelerator.FindDeviceForPointer(dst) && + (!client_->CanUseLocalMemcpy(replica) || + client_->IsHotCacheEnabled())) { if (!client_buffer_allocator_) { LOG(ERROR) << "Client buffer allocator is not provided"; return tl::unexpected(ErrorCode::INVALID_PARAMS); @@ -3527,7 +3529,9 @@ tl::expected RealClient::execute_ranged_read( auto runtime_accelerator = device::GetAcceleratorRegistry().RuntimeAccelerators(); void *dst = static_cast(buffer) + dst_offset; - if (runtime_accelerator.FindDeviceForPointer(dst)) { + auto filtered_qr = FilterQueryResult(query_result, replica, false); + if (runtime_accelerator.FindDeviceForPointer(dst) && + !client_->CanUseLocalMemcpy(replica)) { if (!client_buffer_allocator_) { LOG(ERROR) << "Client buffer allocator is not provided"; return tl::unexpected(ErrorCode::INVALID_PARAMS); @@ -3543,7 +3547,7 @@ tl::expected RealClient::execute_ranged_read( tmp_slices.emplace_back(Slice{tmp_handle.ptr(), size}); auto get_result = - client_->Get(key, query_result, tmp_slices, src_offset); + client_->Get(key, filtered_qr, tmp_slices, src_offset); if (!get_result) { return tl::unexpected(get_result.error()); } @@ -3558,7 +3562,7 @@ tl::expected RealClient::execute_ranged_read( std::vector slices; slices.emplace_back(Slice{dst, size}); - auto get_result = client_->Get(key, query_result, slices, src_offset); + auto get_result = client_->Get(key, filtered_qr, slices, src_offset); if (!get_result) { return tl::unexpected(get_result.error()); } @@ -3644,6 +3648,8 @@ RealClient::get_into_ranges_internal( .first->second; }; + auto runtime_accelerator = + device::GetAcceleratorRegistry().RuntimeAccelerators(); struct ScatterLease { std::chrono::steady_clock::time_point expires_at; std::optional error; @@ -3682,6 +3688,53 @@ RealClient::get_into_ranges_internal( const auto &metadata = metadata_result.value(); if (metadata.replica.is_memory_replica()) { + if (client_->CanUseLocalMemcpy(metadata.replica) && + runtime_accelerator.FindDeviceForPointer(buffers[i])) { + // Planning cache entries may be close to expiry. Renew and + // reselect the replica before copying into device memory. + auto refresh_result = resolve_ranged_read_metadata(keys[j]); + if (!refresh_result) { + std::fill(range_results.begin(), range_results.end(), + tl::unexpected(refresh_result.error())); + continue; + } + std::optional refreshed_metadata; + refreshed_metadata.emplace(std::move(*refresh_result)); + auto lease_refresh_at = + [](const RangedReadMetadata &value) { + const auto now = std::chrono::steady_clock::now(); + return now + + (value.query_result.lease_timeout - now) / 2; + }; + auto refresh_at = lease_refresh_at(*refreshed_metadata); + for (size_t k = 0; k < range_results.size(); ++k) { + // Renew halfway through the remaining lease in long + // batches. + const size_t dst_offset = dst_offsets[k]; + if (dst_offset > capacities[i] || + sizes[k] > capacities[i] - dst_offset) { + continue; + } + if (std::chrono::steady_clock::now() >= refresh_at) { + auto next_refresh_result = + resolve_ranged_read_metadata(keys[j]); + if (!next_refresh_result) { + std::fill(range_results.begin() + k, + range_results.end(), + tl::unexpected( + next_refresh_result.error())); + break; + } + refreshed_metadata.emplace( + std::move(*next_refresh_result)); + refresh_at = lease_refresh_at(*refreshed_metadata); + } + range_results[k] = execute_ranged_read( + keys[j], buffers[i], dst_offset, src_offsets[k], + sizes[k], *refreshed_metadata, false, false); + } + continue; + } const auto &handle = metadata.replica.get_memory_descriptor().buffer_descriptor; auto [lease_it, inserted] = scatter_leases.try_emplace(keys[j]); @@ -4633,58 +4686,22 @@ RealClient::batch_get_into_cuda_ipc_dummy_helper( std::vector> results( requests.size(), tl::unexpected(ErrorCode::INVALID_PARAMS)); - std::vector mappings; - std::vector buffers; - std::vector> all_keys; - std::vector>> all_dst_offsets; - std::vector>> all_src_offsets; - std::vector>> all_sizes; - std::vector buffer_capacities; - std::vector original_indices; - mappings.reserve(requests.size()); - buffers.reserve(requests.size()); - all_keys.reserve(requests.size()); - all_dst_offsets.reserve(requests.size()); - all_src_offsets.reserve(requests.size()); - all_sizes.reserve(requests.size()); - buffer_capacities.reserve(requests.size()); - original_indices.reserve(requests.size()); for (size_t i = 0; i < requests.size(); ++i) { const auto &request = requests[i]; + if (request.size == 0) { + results[i] = 0; + continue; + } auto mapping = device::CudaIpcBufferMapping::Open(request.destination); if (!mapping) { results[i] = tl::unexpected(mapping.error()); continue; } - mappings.push_back(std::move(*mapping)); - buffers.push_back(mappings.back().ptr()); - all_keys.push_back({request.key}); - all_dst_offsets.push_back({{0}}); - all_src_offsets.push_back( - {{static_cast(request.source_offset)}}); - all_sizes.push_back({{static_cast(request.size)}}); - buffer_capacities.push_back(static_cast(request.size)); - original_indices.push_back(i); - } - - if (buffers.empty()) { - return results; - } - - auto range_results = get_into_ranges_internal( - buffers, all_keys, all_dst_offsets, all_src_offsets, all_sizes, - &buffer_capacities, nullptr); - for (size_t i = 0; i < original_indices.size(); ++i) { - if (i < range_results.size() && range_results[i].size() == 1 && - range_results[i][0].size() == 1) { - results[original_indices[i]] = range_results[i][0][0]; - } else { - LOG(ERROR) << "Invalid cuda ipc tensor read result shape for key " - << requests[original_indices[i]].key; - results[original_indices[i]] = - tl::unexpected(ErrorCode::INTERNAL_ERROR); - } + results[i] = get_into_range_internal( + request.key, mapping->ptr(), 0, + static_cast(request.source_offset), + static_cast(request.size), false, false); } return results; } diff --git a/mooncake-store/tests/transfer_task_test.cpp b/mooncake-store/tests/transfer_task_test.cpp index 7ec18f635f..de8bd8d4b8 100644 --- a/mooncake-store/tests/transfer_task_test.cpp +++ b/mooncake-store/tests/transfer_task_test.cpp @@ -5,6 +5,7 @@ #include #include +#include #include #include #include @@ -21,6 +22,36 @@ namespace mooncake { // Test fixture for TransferTask tests +// TODO: Currently, this test does not cover TransferSubmitter and +// TransferEngine integration. Will add more tests in the future. +class ScopedEnvVar { + public: + ScopedEnvVar(const char* name, const char* value) : name_(name) { + if (const char* old_value = std::getenv(name)) { + had_old_value_ = true; + old_value_ = old_value; + } + if (value) { + setenv(name_.c_str(), value, 1); + } else { + unsetenv(name_.c_str()); + } + } + + ~ScopedEnvVar() { + if (had_old_value_) { + setenv(name_.c_str(), old_value_.c_str(), 1); + } else { + unsetenv(name_.c_str()); + } + } + + private: + std::string name_; + bool had_old_value_ = false; + std::string old_value_; +}; + class TransferTaskTest : public ::testing::Test { protected: void SetUp() override { @@ -391,6 +422,38 @@ TEST_F(TransferTaskTest, BatchGetOffloadObjectCopiesPinnedHostToGpu) { EXPECT_EQ(cudaFree(gpu_destination), cudaSuccess); } #endif +TEST_F(TransferTaskTest, CanUseLocalMemcpyRequiresSameProcessEndpoint) { + ScopedEnvVar memcpy_enabled("MC_STORE_MEMCPY", "1"); + + TransferEngine engine(false); + ASSERT_EQ( + engine.init("P2PHANDSHAKE", "127.0.0.1:30991", "127.0.0.1", 30991), 0); + ASSERT_NE(engine.installTransport("tcp", nullptr), nullptr); + + std::shared_ptr storage_backend; + TransferSubmitter submitter(engine, storage_backend, "127.0.0.1:30991"); + + const auto local_endpoint = engine.getLocalIpAndPort(); + ASSERT_FALSE(local_endpoint.empty()); + EXPECT_TRUE(submitter.canUseLocalMemcpy(local_endpoint)); + + EXPECT_FALSE(submitter.canUseLocalMemcpy("127.0.0.1:30992")); + EXPECT_FALSE(submitter.canUseLocalMemcpy("")); +} + +TEST_F(TransferTaskTest, CanUseLocalMemcpyHonorsMemcpyEnv) { + ScopedEnvVar memcpy_disabled("MC_STORE_MEMCPY", "0"); + + TransferEngine engine(false); + ASSERT_EQ( + engine.init("P2PHANDSHAKE", "127.0.0.1:30993", "127.0.0.1", 30993), 0); + ASSERT_NE(engine.installTransport("tcp", nullptr), nullptr); + + std::shared_ptr storage_backend; + TransferSubmitter submitter(engine, storage_backend, "127.0.0.1:30993"); + + EXPECT_FALSE(submitter.canUseLocalMemcpy(engine.getLocalIpAndPort())); +} // Test TransferStrategy enum and stream operator TEST_F(TransferTaskTest, TransferStrategyEnum) { From 9ca10626d95d52182d9a647ba22d1c0483461c9f Mon Sep 17 00:00:00 2001 From: Schatten Date: Tue, 11 Aug 2026 14:11:42 +0800 Subject: [PATCH 029/483] [Store] Add deterministic MasterScenario object lifecycle coverage (#3368) Signed-off-by: Schatten --- mooncake-store/tests/master_scenario.cpp | 4 +- mooncake-store/tests/master_scenario.h | 12 ++ mooncake-store/tests/master_scenario_test.cpp | 14 ++ .../tests/master_service_scenario_test.cpp | 57 +++++++ mooncake-store/tests/master_service_test.cpp | 140 ------------------ 5 files changed, 85 insertions(+), 142 deletions(-) diff --git a/mooncake-store/tests/master_scenario.cpp b/mooncake-store/tests/master_scenario.cpp index 0962ba1ca7..da1032d570 100644 --- a/mooncake-store/tests/master_scenario.cpp +++ b/mooncake-store/tests/master_scenario.cpp @@ -92,7 +92,7 @@ MasterScenario& MasterScenario::WhenPutStart(PutStartActionData action) { } ReplicateConfig config; - config.replica_num = 1; + config.replica_num = action.requested_replica_count; const auto result = service_->PutStart(ActorId(action.actor), action.key, TenantId::Default(), action.size, config); @@ -108,7 +108,7 @@ MasterScenario& MasterScenario::WhenUpsertStart(UpsertStartActionData action) { } ReplicateConfig config; - config.replica_num = 1; + config.replica_num = action.requested_replica_count; const auto result = service_->UpsertStart(ActorId(action.actor), action.key, TenantId::Default(), action.size, config); diff --git a/mooncake-store/tests/master_scenario.h b/mooncake-store/tests/master_scenario.h index f6948b7078..3a1e859b5b 100644 --- a/mooncake-store/tests/master_scenario.h +++ b/mooncake-store/tests/master_scenario.h @@ -51,6 +51,7 @@ struct PutStartActionData { std::string key; uint64_t size; std::string actor{"default"}; + size_t requested_replica_count{1}; std::optional expected_error{}; std::optional expected_replica_count{}; std::optional expected_replica_status{}; @@ -71,6 +72,11 @@ struct PutStartAction : PutStartActionData { return *this; } + PutStartAction& Replicas(size_t value) { + requested_replica_count = value; + return *this; + } + auto ExpectError(ErrorCode value) const requires(expectation == PutStartExpectation::UNSPECIFIED) { @@ -126,6 +132,7 @@ struct UpsertStartActionData { std::string key; uint64_t size; std::string actor{"default"}; + size_t requested_replica_count{1}; std::optional expected_error{}; std::optional expected_replica_count{}; std::optional expected_replica_status{}; @@ -146,6 +153,11 @@ struct UpsertStartAction : UpsertStartActionData { return *this; } + UpsertStartAction& Replicas(size_t value) { + requested_replica_count = value; + return *this; + } + auto ExpectError(ErrorCode value) const requires(expectation == UpsertStartExpectation::UNSPECIFIED) { diff --git a/mooncake-store/tests/master_scenario_test.cpp b/mooncake-store/tests/master_scenario_test.cpp index 4fc5cb645f..0edb9034ad 100644 --- a/mooncake-store/tests/master_scenario_test.cpp +++ b/mooncake-store/tests/master_scenario_test.cpp @@ -92,6 +92,20 @@ static_assert(SupportsThen); } // namespace +TEST(MasterScenarioContractTest, HonorsRequestedPutStartReplicaCount) { + MasterScenario("requested put start replica count") + .Given(MemoryNode("memory-1")) + .Given(MemoryNode("memory-2")) + .When(PutStart("key", 1_KB).Replicas(2).ExpectReplicas(2)); +} + +TEST(MasterScenarioContractTest, HonorsRequestedUpsertStartReplicaCount) { + MasterScenario("requested upsert start replica count") + .Given(MemoryNode("memory-1")) + .Given(MemoryNode("memory-2")) + .When(UpsertStart("key", 1_KB).Replicas(2).ExpectReplicas(2)); +} + TEST(MasterScenarioContractTest, ReportsUnexpectedActionError) { EXPECT_NONFATAL_FAILURE(MasterScenario("unexpected action error") .Given(MemoryNode("memory")) diff --git a/mooncake-store/tests/master_service_scenario_test.cpp b/mooncake-store/tests/master_service_scenario_test.cpp index 822925ec14..7ffb27573f 100644 --- a/mooncake-store/tests/master_service_scenario_test.cpp +++ b/mooncake-store/tests/master_service_scenario_test.cpp @@ -1,5 +1,7 @@ #include "master_scenario.h" +#include + #include namespace mooncake::test { @@ -26,6 +28,61 @@ TEST(MasterServiceTest, PutStartEndFlow) { .HasCompleteReplicas(1)); } +TEST(MasterServiceTest, PutLifecycleForEveryReplicaCount) { + constexpr std::array kReplicaCounts{1, 2, 3, 4, 5}; + MasterScenario scenario( + "put lifecycle for replica counts one through five"); + for (size_t index : kReplicaCounts) { + scenario.Given(MemoryNode("memory-" + std::to_string(index))); + } + for (size_t replica_count : kReplicaCounts) { + const std::string key = "key-" + std::to_string(replica_count); + scenario + .When(PutStart(key, 1_KB) + .Replicas(replica_count) + .ExpectReplicas(replica_count) + .ExpectStatus(ReplicaStatus::PROCESSING)) + .Then(Object(key).IsNotReady()) + .When(Remove(key).ExpectError(ErrorCode::REPLICA_IS_NOT_READY)) + .When(PutEnd(key)) + .Then(Object(key) + .IsReadable() + .HasReplicas(replica_count) + .HasCompleteReplicas(replica_count)); + } +} + +TEST(MasterServiceTest, GetReplicaListDistinguishesMissingAndReadable) { + MasterScenario("get replica list distinguishes missing and readable") + .Given(MemoryNode("memory")) + .Then(Object("missing").DoesNotExist()) + .When(PutStart("key", 1_KB)) + .When(PutEnd("key")) + .Then(Object("key").IsReadable().HasReplicas(1)); +} + +TEST(MasterServiceTest, RemoveObjectAndRejectMissingObject) { + MasterScenario("remove object and reject a missing object") + .Given(MemoryNode("memory")) + .When(PutStart("key", 1_KB)) + .When(PutEnd("key")) + .When(Remove("key")) + .Then(Object("key").DoesNotExist()) + .When(Remove("missing").ExpectError(ErrorCode::OBJECT_NOT_FOUND)); +} + +TEST(MasterServiceTest, RepeatedPutAndRemoveIsDeterministic) { + MasterScenario scenario("repeated put and remove with fixed keys"); + scenario.Given(MemoryNode("memory")); + for (int index = 0; index < 10; ++index) { + const std::string key = "key-" + std::to_string(index); + scenario.When(PutStart(key, 1_KB)) + .When(PutEnd(key)) + .When(Remove(key)) + .Then(Object(key).DoesNotExist()); + } +} + TEST(MasterServiceTest, UpsertNewKey) { MasterScenario("upsert creates a new object") .Given(MemoryNode("memory")) diff --git a/mooncake-store/tests/master_service_test.cpp b/mooncake-store/tests/master_service_test.cpp index 920d86e1c0..02b072782a 100644 --- a/mooncake-store/tests/master_service_test.cpp +++ b/mooncake-store/tests/master_service_test.cpp @@ -2406,55 +2406,6 @@ TEST_F(MasterServiceTest, ExplicitPreferredSegmentFallsBackToLocalFirst) { "segment_host1"); } -TEST_F(MasterServiceTest, RandomPutStartEndFlow) { - std::unique_ptr service_(new MasterService()); - const UUID client_id = generate_uuid(); - - // Mount 5 segments, each 16MB - constexpr size_t kBaseAddr = 0x300000000; - constexpr size_t kSegmentSize = 1024 * 1024 * 16; // 16MB - for (int i = 0; i < 5; ++i) { - [[maybe_unused]] const auto context = PrepareSimpleSegment( - *service_, "segment_" + std::to_string(i), - kBaseAddr + static_cast(i) * kSegmentSize, kSegmentSize); - } - - // Test PutStart - std::string key = "test_key"; - uint64_t value_length = 1024; - ReplicateConfig config; - std::random_device rd; - std::mt19937 gen(rd()); - std::uniform_int_distribution<> dis(1, 5); - int random_number = dis(gen); - config.replica_num = random_number; - auto put_start_result = service_->PutStart( - client_id, key, TenantId::Default(), value_length, config); - EXPECT_TRUE(put_start_result.has_value()); - replica_list = put_start_result.value(); - EXPECT_FALSE(replica_list.empty()); - EXPECT_EQ(ReplicaStatus::PROCESSING, replica_list[0].status); - // During put, Get/Remove should fail - auto get_result = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_FALSE(get_result.has_value()); - EXPECT_EQ(ErrorCode::REPLICA_IS_NOT_READY, get_result.error()); - auto remove_result = service_->Remove(key, TenantId::Default()); - EXPECT_FALSE(remove_result.has_value()); - EXPECT_EQ(ErrorCode::REPLICA_IS_NOT_READY, remove_result.error()); - // Test PutEnd - auto put_end_result = service_->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - EXPECT_TRUE(put_end_result.has_value()); - // Verify replica list after PutEnd - auto get_result2 = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_TRUE(get_result2.has_value()); - replica_list = get_result2.value().replicas; - EXPECT_EQ(random_number, replica_list.size()); - for (int i = 0; i < random_number; ++i) { - EXPECT_EQ(ReplicaStatus::COMPLETE, replica_list[i].status); - } -} - TEST_F(MasterServiceTest, GetReplicaListByRegex) { const uint64_t kv_lease_ttl = 50; auto service_config = MasterServiceConfig::builder() @@ -2641,97 +2592,6 @@ TEST_F(MasterServiceTest, GetReplicaListByRegexComplex) { } } -TEST_F(MasterServiceTest, GetReplicaList) { - std::unique_ptr service_(new MasterService()); - const UUID client_id = generate_uuid(); - // Test getting non-existent key - auto get_result = - service_->GetReplicaList("non_existent", TenantId::Default()); - EXPECT_FALSE(get_result.has_value()); - EXPECT_EQ(ErrorCode::OBJECT_NOT_FOUND, get_result.error()); - - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - - std::string key = "test_key"; - uint64_t value_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - auto put_start_result = service_->PutStart( - client_id, key, TenantId::Default(), value_length, config); - ASSERT_TRUE(put_start_result.has_value()); - auto put_end_result = service_->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end_result.has_value()); - - // Test getting existing key - auto get_result2 = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_TRUE(get_result2.has_value()); - auto replica_list_local = get_result2.value().replicas; - EXPECT_FALSE(replica_list_local.empty()); -} - -TEST_F(MasterServiceTest, RemoveObject) { - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - - std::string key = "test_key"; - uint64_t value_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - auto put_start_result = service_->PutStart( - client_id, key, TenantId::Default(), value_length, config); - ASSERT_TRUE(put_start_result.has_value()); - auto put_end_result = service_->PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY); - ASSERT_TRUE(put_end_result.has_value()); - - // Test removing the object - auto remove_result = service_->Remove(key, TenantId::Default()); - EXPECT_TRUE(remove_result.has_value()); - - // Verify object is removed - auto get_result = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_FALSE(get_result.has_value()); - EXPECT_EQ(ErrorCode::OBJECT_NOT_FOUND, get_result.error()); - - // Test removing non-existent object - auto remove_result2 = service_->Remove("non_existent", TenantId::Default()); - EXPECT_FALSE(remove_result2.has_value()); - EXPECT_EQ(ErrorCode::OBJECT_NOT_FOUND, remove_result2.error()); -} - -TEST_F(MasterServiceTest, RandomRemoveObject) { - std::unique_ptr service_(new MasterService()); - [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); - const UUID client_id = generate_uuid(); - int times = 10; - std::random_device rd; - std::mt19937 gen(rd()); - std::uniform_int_distribution<> dis(1, 1000); - while (times--) { - std::string key = "test_key" + std::to_string(dis(gen)); - uint64_t value_length = 1024; - ReplicateConfig config; - config.replica_num = 1; - auto put_start_result = service_->PutStart( - client_id, key, TenantId::Default(), value_length, config); - ASSERT_TRUE(put_start_result.has_value()); - auto put_end_result = service_->PutEnd( - client_id, key, TenantId::Default(), ReplicaType::MEMORY); - ASSERT_TRUE(put_end_result.has_value()); - - // Test removing the object - auto remove_result = service_->Remove(key, TenantId::Default()); - EXPECT_TRUE(remove_result.has_value()); - - // Verify object is removed - auto get_result = service_->GetReplicaList(key, TenantId::Default()); - EXPECT_FALSE(get_result.has_value()); - EXPECT_EQ(ErrorCode::OBJECT_NOT_FOUND, get_result.error()); - } -} - TEST_F(MasterServiceTest, RemoveByRegex) { const uint64_t kv_lease_ttl = 50; auto service_config = MasterServiceConfig::builder() From bb9b78379772df9ea90c70635b0315423696a55b Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Tue, 11 Aug 2026 14:20:08 +0800 Subject: [PATCH 030/483] [TransferEngine] Fix graceful shutdown test hangs in USE_ETCD builds (#3351) * [TransferEngine] Fix graceful shutdown signal tests * [TransferEngine] Improve graceful shutdown test diagnostics --------- Co-authored-by: Yuchen Kou --- mooncake-transfer-engine/tests/CMakeLists.txt | 3 +- .../tests/graceful_shutdown_test.cpp | 295 +++++++++++++----- 2 files changed, 218 insertions(+), 80 deletions(-) diff --git a/mooncake-transfer-engine/tests/CMakeLists.txt b/mooncake-transfer-engine/tests/CMakeLists.txt index 6935fd7e96..545acab27c 100644 --- a/mooncake-transfer-engine/tests/CMakeLists.txt +++ b/mooncake-transfer-engine/tests/CMakeLists.txt @@ -361,8 +361,7 @@ if(ENABLE_MULTI_PROTOCOL) endif() add_executable(graceful_shutdown_test ${WORKSPACE}/graceful_shutdown_test.cpp) -target_link_libraries(graceful_shutdown_test PUBLIC transfer_engine gtest - gtest_main) +target_link_libraries(graceful_shutdown_test PUBLIC transfer_engine gtest) add_test(NAME graceful_shutdown_test COMMAND graceful_shutdown_test) add_executable(show_links_test ${WORKSPACE}/show_links_test.cpp) diff --git a/mooncake-transfer-engine/tests/graceful_shutdown_test.cpp b/mooncake-transfer-engine/tests/graceful_shutdown_test.cpp index f0741c82bf..8d8c320002 100644 --- a/mooncake-transfer-engine/tests/graceful_shutdown_test.cpp +++ b/mooncake-transfer-engine/tests/graceful_shutdown_test.cpp @@ -13,72 +13,228 @@ // limitations under the License. #include +#include #include +#include #include #include +#include +#include +#include +#include +#include +#include #include #include "transfer_engine.h" +extern char** environ; + using namespace mooncake; namespace { -void waitChildWithTimeout(pid_t pid, int* status) { - for (int i = 0; i < 50; ++i) { - pid_t ret = waitpid(pid, status, WNOHANG); - ASSERT_NE(ret, -1) << "waitpid() failed"; - if (ret == pid) return; - usleep(100000); +constexpr char kChildFlag[] = "--graceful-shutdown-child"; +constexpr int kChildTimeoutMs = 5000; + +enum class ChildMode { kSingle, kDestroyed, kMultiple }; + +const char* childModeName(ChildMode mode) { + switch (mode) { + case ChildMode::kSingle: + return "single"; + case ChildMode::kDestroyed: + return "destroyed"; + case ChildMode::kMultiple: + return "multiple"; } - kill(pid, SIGKILL); - waitpid(pid, status, 0); - FAIL() << "child did not exit before timeout"; + return "unknown"; } -} // namespace +bool parseChildMode(const char* value, ChildMode* mode) { + if (strcmp(value, "single") == 0) { + *mode = ChildMode::kSingle; + } else if (strcmp(value, "destroyed") == 0) { + *mode = ChildMode::kDestroyed; + } else if (strcmp(value, "multiple") == 0) { + *mode = ChildMode::kMultiple; + } else { + return false; + } + return true; +} -TEST(GracefulShutdownTest, SigtermTriggersCleanExit) { - pid_t pid = fork(); - ASSERT_NE(pid, -1) << "fork() failed"; +int runShutdownChild(ChildMode mode, int ready_fd) { + std::unique_ptr engine1; + std::unique_ptr engine2; - if (pid == 0) { + if (mode == ChildMode::kDestroyed) { auto engine = std::make_unique(false); engine->enableGracefulShutdown(); - pause(); - _exit(99); + } else { + engine1 = std::make_unique(false); + engine1->enableGracefulShutdown(); + if (mode == ChildMode::kMultiple) { + engine2 = std::make_unique(false); + engine2->enableGracefulShutdown(); + } + } + + const char ready = '1'; + ssize_t bytes_written; + do { + bytes_written = write(ready_fd, &ready, sizeof(ready)); + } while (bytes_written < 0 && errno == EINTR); + if (bytes_written != static_cast(sizeof(ready))) return 111; + close(ready_fd); + for (;;) pause(); +} + +void killAndReap(pid_t pid) { + if (kill(pid, SIGKILL) != 0 && errno != ESRCH) { + ADD_FAILURE() << "kill(SIGKILL) failed: " << strerror(errno); + } + + int status = 0; + pid_t ret; + do { + ret = waitpid(pid, &status, 0); + } while (ret < 0 && errno == EINTR); + if (ret < 0) { + ADD_FAILURE() << "waitpid() failed while reaping child: " + << strerror(errno); + } +} + +bool waitChildWithTimeout(pid_t pid, int* status) { + auto deadline = std::chrono::steady_clock::now() + + std::chrono::milliseconds(kChildTimeoutMs); + for (;;) { + pid_t ret = waitpid(pid, status, WNOHANG); + if (ret == pid) return true; + if (ret < 0) { + if (errno == EINTR) continue; + int error = errno; + if (error != ECHILD) killAndReap(pid); + ADD_FAILURE() << "waitpid() failed: " << strerror(error); + return false; + } + if (std::chrono::steady_clock::now() >= deadline) break; + usleep(100000); + } + + killAndReap(pid); + ADD_FAILURE() << "child did not exit before timeout"; + return false; +} + +bool waitForChildReady(int fd) { + pollfd pfd{fd, POLLIN, 0}; + auto deadline = std::chrono::steady_clock::now() + + std::chrono::milliseconds(kChildTimeoutMs); + int ret; + for (;;) { + auto remaining = std::chrono::duration_cast( + deadline - std::chrono::steady_clock::now()) + .count(); + if (remaining <= 0) { + ret = 0; + break; + } + ret = poll(&pfd, 1, static_cast(remaining)); + if (ret >= 0 || errno != EINTR) break; + } + + if (ret == 0) { + ADD_FAILURE() << "child did not report readiness before timeout"; + return false; + } + if (ret < 0) { + ADD_FAILURE() << "poll() failed: " << strerror(errno); + return false; + } + + char ready = 0; + ssize_t bytes_read; + do { + bytes_read = read(fd, &ready, sizeof(ready)); + } while (bytes_read < 0 && errno == EINTR); + if (bytes_read == 0) { + ADD_FAILURE() << "child exited before reporting readiness"; + return false; + } + if (bytes_read < 0) { + ADD_FAILURE() << "read() failed while waiting for readiness: " + << strerror(errno); + return false; + } + if (ready != '1') { + ADD_FAILURE() << "unexpected readiness marker: " + << static_cast(static_cast(ready)); + return false; + } + return true; +} + +pid_t spawnReadyChild(ChildMode mode) { + int ready_pipe[2]; + if (pipe(ready_pipe) != 0) { + ADD_FAILURE() << "pipe() failed: " << strerror(errno); + return -1; + } + + char ready_fd[32]; + snprintf(ready_fd, sizeof(ready_fd), "%d", ready_pipe[1]); + char* child_argv[] = { + const_cast("/proc/self/exe"), const_cast(kChildFlag), + const_cast(childModeName(mode)), ready_fd, nullptr}; + + pid_t pid = -1; + int spawn_error = posix_spawn(&pid, "/proc/self/exe", nullptr, nullptr, + child_argv, environ); + close(ready_pipe[1]); + if (spawn_error != 0) { + close(ready_pipe[0]); + ADD_FAILURE() << "posix_spawn() failed: " << strerror(spawn_error); + return -1; + } + + bool ready = waitForChildReady(ready_pipe[0]); + close(ready_pipe[0]); + if (!ready) { + killAndReap(pid); + return -1; } + return pid; +} - usleep(100000); - kill(pid, SIGTERM); +void expectGracefulExit(ChildMode mode, int signo) { + pid_t pid = spawnReadyChild(mode); + ASSERT_GT(pid, 0); - int status; - waitChildWithTimeout(pid, &status); + if (kill(pid, signo) != 0) { + int error = errno; + killAndReap(pid); + FAIL() << "kill(" << signo << ") failed: " << strerror(error); + } + + int status = 0; + ASSERT_TRUE(waitChildWithTimeout(pid, &status)); ASSERT_TRUE(WIFEXITED(status)) << "Child did not exit normally (signaled: " << WIFSIGNALED(status) << ")"; - EXPECT_EQ(WEXITSTATUS(status), 128 + SIGTERM); + EXPECT_EQ(WEXITSTATUS(status), 128 + signo); } -TEST(GracefulShutdownTest, SigintTriggersCleanExit) { - pid_t pid = fork(); - ASSERT_NE(pid, -1) << "fork() failed"; - - if (pid == 0) { - auto engine = std::make_unique(false); - engine->enableGracefulShutdown(); - pause(); - _exit(99); - } +} // namespace - usleep(100000); - kill(pid, SIGINT); +TEST(GracefulShutdownTest, SigtermTriggersCleanExit) { + expectGracefulExit(ChildMode::kSingle, SIGTERM); +} - int status; - waitChildWithTimeout(pid, &status); - ASSERT_TRUE(WIFEXITED(status)) << "Child did not exit normally"; - EXPECT_EQ(WEXITSTATUS(status), 128 + SIGINT); +TEST(GracefulShutdownTest, SigintTriggersCleanExit) { + expectGracefulExit(ChildMode::kSingle, SIGINT); } TEST(GracefulShutdownTest, IdempotentEnable) { @@ -89,25 +245,7 @@ TEST(GracefulShutdownTest, IdempotentEnable) { } TEST(GracefulShutdownTest, EngineDestroyedBeforeSignal) { - pid_t pid = fork(); - ASSERT_NE(pid, -1) << "fork() failed"; - - if (pid == 0) { - { - auto engine = std::make_unique(false); - engine->enableGracefulShutdown(); - } - pause(); - _exit(99); - } - - usleep(100000); - kill(pid, SIGTERM); - - int status; - waitChildWithTimeout(pid, &status); - ASSERT_TRUE(WIFEXITED(status)); - EXPECT_EQ(WEXITSTATUS(status), 128 + SIGTERM); + expectGracefulExit(ChildMode::kDestroyed, SIGTERM); } TEST(GracefulShutdownTest, ForkAfterInstallDoesNotHangChildSignal) { @@ -118,15 +256,17 @@ TEST(GracefulShutdownTest, ForkAfterInstallDoesNotHangChildSignal) { ASSERT_NE(pid, -1) << "fork() failed"; if (pid == 0) { - pause(); - _exit(99); + for (;;) pause(); } - usleep(100000); - kill(pid, SIGTERM); + if (kill(pid, SIGTERM) != 0) { + int error = errno; + killAndReap(pid); + FAIL() << "kill(SIGTERM) failed: " << strerror(error); + } - int status; - waitChildWithTimeout(pid, &status); + int status = 0; + ASSERT_TRUE(waitChildWithTimeout(pid, &status)); ASSERT_TRUE(WIFEXITED(status)) << "Child did not exit normally (signaled: " << WIFSIGNALED(status) << ")"; @@ -134,23 +274,22 @@ TEST(GracefulShutdownTest, ForkAfterInstallDoesNotHangChildSignal) { } TEST(GracefulShutdownTest, MultipleEngines) { - pid_t pid = fork(); - ASSERT_NE(pid, -1) << "fork() failed"; + expectGracefulExit(ChildMode::kMultiple, SIGTERM); +} - if (pid == 0) { - auto engine1 = std::make_unique(false); - auto engine2 = std::make_unique(false); - engine1->enableGracefulShutdown(); - engine2->enableGracefulShutdown(); - pause(); - _exit(99); +int main(int argc, char** argv) { + if (argc == 4 && strcmp(argv[1], kChildFlag) == 0) { + ChildMode mode; + char* end = nullptr; + errno = 0; + long ready_fd = strtol(argv[3], &end, 10); + if (!parseChildMode(argv[2], &mode) || errno != 0 || *end != '\0' || + ready_fd < 0 || ready_fd > INT_MAX) { + return 2; + } + return runShutdownChild(mode, static_cast(ready_fd)); } - usleep(100000); - kill(pid, SIGTERM); - - int status; - waitChildWithTimeout(pid, &status); - ASSERT_TRUE(WIFEXITED(status)); - EXPECT_EQ(WEXITSTATUS(status), 128 + SIGTERM); + testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); } From b82b9f86ceb67bb809339b9088085f34329683bb Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Tue, 11 Aug 2026 14:24:22 +0800 Subject: [PATCH 031/483] [Store] Fix race in HA durable-finalization tests (#3332) (#3333) Co-authored-by: Yuchen Kou --- .../include/ha/oplog/ordered_oplog_writer.h | 9 +- mooncake-store/include/master_service.h | 5 + mooncake-store/src/master_service.cpp | 28 ++- .../tests/ha/master_service_ha_test.cpp | 226 ++++++++++++------ 4 files changed, 191 insertions(+), 77 deletions(-) diff --git a/mooncake-store/include/ha/oplog/ordered_oplog_writer.h b/mooncake-store/include/ha/oplog/ordered_oplog_writer.h index 1d4387d51a..49f54afa02 100644 --- a/mooncake-store/include/ha/oplog/ordered_oplog_writer.h +++ b/mooncake-store/include/ha/oplog/ordered_oplog_writer.h @@ -54,18 +54,17 @@ class OrderedOpLogWriter { OrderedOpLogWriter(OrderedOpLogWriterConfig config, WriteBatchFn write_batch); - ~OrderedOpLogWriter(); + virtual ~OrderedOpLogWriter(); tl::expected Reserve(); - tl::expected Commit(Reservation&& reservation, - OpLogEntry entry, - DurableCallback callback); + virtual tl::expected Commit( + Reservation&& reservation, OpLogEntry entry, DurableCallback callback); void Abort(Reservation&& reservation); bool IsAccepting() const; ErrorCode LastError() const; void Start(); - void Stop(); + virtual void Stop(); private: struct Impl; diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index d341993059..cd2d357100 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -136,6 +136,9 @@ class MasterService { std::function; using DurableFinalizeCallback = std::function; + using BatchOpLogWriterFactory = + std::function( + OrderedOpLogWriterConfig, OrderedOpLogWriter::WriteBatchFn)>; MasterService(); MasterService(const MasterServiceConfig& config); @@ -158,6 +161,7 @@ class MasterService { ErrorCode SetBatchOpLogBackendForTesting( std::shared_ptr backend); + void SetBatchOpLogWriterFactoryForTesting(BatchOpLogWriterFactory factory); /** * @brief Test-only wrapper around BatchEvict / NoFBatchEvict so that @@ -2491,6 +2495,7 @@ class MasterService { std::shared_ptr batch_oplog_kv_backend_; std::unique_ptr batch_oplog_storage_; std::unique_ptr ordered_oplog_writer_; + BatchOpLogWriterFactory batch_oplog_writer_factory_; // OpLog publishing helpers std::string SerializeMetadataForOpLog(const ObjectMetadata& metadata) const; diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index 05494f793c..c61a3af482 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -201,7 +201,13 @@ MasterService::MasterService(const MasterServiceConfig& config) put_start_release_timeout_sec_(config.put_start_release_timeout_sec), offloading_queue_limit_(config.offloading_queue_limit), offload_cap_ratio_(config.offload_cap_ratio), - task_manager_(config.task_manager_config) { + task_manager_(config.task_manager_config), + batch_oplog_writer_factory_( + [](OrderedOpLogWriterConfig writer_config, + OrderedOpLogWriter::WriteBatchFn write_batch) { + return std::make_unique( + std::move(writer_config), std::move(write_batch)); + }) { if (default_kv_soft_pin_ttl_ > max_kv_soft_pin_ttl_) { LOG(ERROR) << "Invalid soft-pin TTL configuration: default=" << default_kv_soft_pin_ttl_ @@ -601,6 +607,13 @@ ErrorCode MasterService::SetBatchOpLogBackendForTesting( return InitializeBatchOpLogWriter(std::move(backend)); } +void MasterService::SetBatchOpLogWriterFactoryForTesting( + BatchOpLogWriterFactory factory) { + assert(factory); + assert(!ordered_oplog_writer_); + batch_oplog_writer_factory_ = std::move(factory); +} + void MasterService::RunBatchEvictForTesting(double evict_ratio_target, double evict_ratio_lowerbound) { BatchEvict(evict_ratio_target, evict_ratio_lowerbound); @@ -11362,12 +11375,17 @@ ErrorCode MasterService::InitializeBatchOpLogWriter( writer_config.max_entries_per_batch = oplog_batch_max_entries_; writer_config.initial_durable_prefix = durable_prefix; OpLogBatchStorage* storage_ptr = storage.get(); - auto writer = std::make_unique( - writer_config, [storage_ptr](const OpLogBatchRecord& batch, - const DurablePrefix& expected_prefix) { + OrderedOpLogWriter::WriteBatchFn write_batch = + [storage_ptr](const OpLogBatchRecord& batch, + const DurablePrefix& expected_prefix) { return storage_ptr->WriteBatchAndAdvancePrefix(batch, expected_prefix); - }); + }; + auto writer = + batch_oplog_writer_factory_(writer_config, std::move(write_batch)); + if (!writer) { + return ErrorCode::INVALID_PARAMS; + } if (!writer->IsAccepting()) { return writer->LastError(); } diff --git a/mooncake-store/tests/ha/master_service_ha_test.cpp b/mooncake-store/tests/ha/master_service_ha_test.cpp index a898480109..a617cc2880 100644 --- a/mooncake-store/tests/ha/master_service_ha_test.cpp +++ b/mooncake-store/tests/ha/master_service_ha_test.cpp @@ -3,6 +3,7 @@ #include #include +#include #include #include #include @@ -149,6 +150,74 @@ class FailingBatchHaKvBackend : public FakeBatchHaKvBackend { size_t txn_calls_{0}; }; +class GatedOrderedOpLogWriter : public OrderedOpLogWriter { + public: + GatedOrderedOpLogWriter(OrderedOpLogWriterConfig config, + WriteBatchFn write_batch) + : OrderedOpLogWriter(std::move(config), std::move(write_batch)) {} + + ~GatedOrderedOpLogWriter() override { Stop(); } + + tl::expected Commit( + Reservation&& reservation, OpLogEntry entry, + DurableCallback callback) override { + return OrderedOpLogWriter::Commit( + std::move(reservation), std::move(entry), + [this, callback = std::move(callback)](const OpLogEntry& durable) { + { + std::unique_lock lock(mutex_); + cv_.wait(lock, [&] { + return stopping_ || + durable.sequence_id <= released_through_; + }); + } + if (callback) { + callback(durable); + } + { + std::lock_guard lock(mutex_); + completed_through_ = durable.sequence_id; + } + cv_.notify_all(); + }); + } + + bool PauseCallbacksAfter( + uint64_t sequence_id, + std::chrono::milliseconds timeout = std::chrono::seconds(1)) { + std::unique_lock lock(mutex_); + released_through_ = sequence_id; + return cv_.wait_for(lock, timeout, + [&] { return completed_through_ >= sequence_id; }); + } + + bool RunCallbacksThrough( + uint64_t sequence_id, + std::chrono::milliseconds timeout = std::chrono::seconds(1)) { + std::unique_lock lock(mutex_); + released_through_ = std::max(released_through_, sequence_id); + cv_.notify_all(); + return cv_.wait_for(lock, timeout, + [&] { return completed_through_ >= sequence_id; }); + } + + void Stop() override { + { + std::lock_guard lock(mutex_); + stopping_ = true; + } + cv_.notify_all(); + OrderedOpLogWriter::Stop(); + } + + private: + std::mutex mutex_; + std::condition_variable cv_; + uint64_t released_through_{UINT64_MAX}; + uint64_t completed_through_{0}; + bool stopping_{false}; +}; + class MasterServiceHATest : public ::testing::Test { protected: static void SetUpTestSuite() { @@ -320,6 +389,26 @@ class MasterServiceHATest : public ::testing::Test { ASSERT_EQ(ErrorCode::OK, read_err); } + static GatedOrderedOpLogWriter* InstallGatedWriter( + MasterService& service, std::shared_ptr backend) { + service.SetBatchOpLogWriterFactoryForTesting( + [](OrderedOpLogWriterConfig config, + OrderedOpLogWriter::WriteBatchFn write_batch) { + return std::make_unique( + std::move(config), std::move(write_batch)); + }); + EXPECT_EQ(ErrorCode::OK, + service.SetBatchOpLogBackendForTesting(std::move(backend))); + return static_cast( + service.ordered_oplog_writer_.get()); + } + + static uint64_t TenantUsedBytes(MasterService& service) { + auto snapshot = service.GetTenantQuotaSnapshot(kDefaultTenant); + EXPECT_TRUE(snapshot.has_value()); + return snapshot ? snapshot->used_bytes : 0; + } + void ReadRemoveBatchEventually(OpLogBatchStorage& storage, uint64_t first_batch_id, const std::string& key, @@ -1699,7 +1788,7 @@ TEST_F(MasterServiceBatchRecordE2ETest, WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "batch_e2e_remove_finalize_segment"); @@ -1711,6 +1800,7 @@ TEST_F(MasterServiceBatchRecordE2ETest, PutObjectOnSegment(service, mounted.client_id, key, "batch_e2e_remove_finalize_segment"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); backend->BlockTxn(); ASSERT_TRUE( @@ -1725,6 +1815,9 @@ TEST_F(MasterServiceBatchRecordE2ETest, backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); auto removed = service.ExistKey(key, kDefaultTenant); ASSERT_TRUE(removed.has_value()); @@ -2780,7 +2873,7 @@ TEST_F(MasterServiceHATest, RemoveHidesBeforeDurableAndReleasesAfterFinalize) { WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "batch_remove_finalize_seg"); OpLogBatchStorage storage(cluster_id, *backend); @@ -2791,6 +2884,7 @@ TEST_F(MasterServiceHATest, RemoveHidesBeforeDurableAndReleasesAfterFinalize) { PutObjectOnSegment(service, mounted.client_id, key, "batch_remove_finalize_seg"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); backend->BlockTxn(); ASSERT_TRUE( @@ -2806,15 +2900,14 @@ TEST_F(MasterServiceHATest, RemoveHidesBeforeDurableAndReleasesAfterFinalize) { backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = "after_remove_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_remove_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, BatchRemoveWritesBatchRecordOpLog) { @@ -2868,7 +2961,7 @@ TEST_F(MasterServiceHATest, BatchRemoveFinalizesEachObjectAfterDurable) { WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "batch_remove_finalize_segment"); @@ -2880,6 +2973,7 @@ TEST_F(MasterServiceHATest, BatchRemoveFinalizesEachObjectAfterDurable) { PutObjectOnSegment(service, mounted.client_id, key, "batch_remove_finalize_segment"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); backend->BlockTxn(); auto results = service.BatchRemove({key}, kDefaultTenant, /*force=*/true); @@ -2896,16 +2990,14 @@ TEST_F(MasterServiceHATest, BatchRemoveFinalizesEachObjectAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = - "after_batch_remove_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_batch_remove_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, RemoveAllWritesBatchRecordOpLog) { @@ -2957,7 +3049,7 @@ TEST_F(MasterServiceHATest, RemoveAllFinalizesAfterDurable) { WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "remove_all_finalize_segment"); OpLogBatchStorage storage(cluster_id, *backend); @@ -2968,6 +3060,7 @@ TEST_F(MasterServiceHATest, RemoveAllFinalizesAfterDurable) { PutObjectOnSegment(service, mounted.client_id, key, "remove_all_finalize_segment"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); backend->BlockTxn(); EXPECT_EQ(1, service.RemoveAll(kDefaultTenant, /*force=*/true)); @@ -2982,15 +3075,14 @@ TEST_F(MasterServiceHATest, RemoveAllFinalizesAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = "after_remove_all_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_remove_all_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, BatchReplicaClearAllWritesBatchRecordOpLog) { @@ -3095,7 +3187,7 @@ TEST_F(MasterServiceHATest, BatchReplicaClearAllReleasesAfterDurable) { WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "clear_all_finalize_segment"); OpLogBatchStorage storage(cluster_id, *backend); @@ -3106,6 +3198,7 @@ TEST_F(MasterServiceHATest, BatchReplicaClearAllReleasesAfterDurable) { PutObjectOnSegment(service, mounted.client_id, key, "clear_all_finalize_segment"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); std::this_thread::sleep_for(std::chrono::milliseconds(60)); backend->BlockTxn(); @@ -3124,15 +3217,14 @@ TEST_F(MasterServiceHATest, BatchReplicaClearAllReleasesAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = "after_clear_all_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_clear_all_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, BatchReplicaClearSegmentReleasesAfterDurable) { @@ -3151,7 +3243,7 @@ TEST_F(MasterServiceHATest, BatchReplicaClearSegmentReleasesAfterDurable) { WriteTenantPolicyFile({{kDefaultTenant.value(), 2048}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "clear_finalize_seg1"); OpLogBatchStorage storage(cluster_id, *backend); @@ -3184,6 +3276,7 @@ TEST_F(MasterServiceHATest, BatchReplicaClearSegmentReleasesAfterDurable) { ASSERT_TRUE( service.CopyEnd(mounted.client_id, key, kDefaultTenant).has_value()); ReadBatchEventually(storage, 4, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); std::this_thread::sleep_for(std::chrono::milliseconds(60)); backend->BlockTxn(); @@ -3202,16 +3295,14 @@ TEST_F(MasterServiceHATest, BatchReplicaClearSegmentReleasesAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 5, batch); + EXPECT_EQ(2048, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(1024, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = - "after_clear_segment_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_clear_segment_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, BatchEvictWritesBatchRecordOpLog) { @@ -3264,7 +3355,7 @@ TEST_F(MasterServiceHATest, BatchEvictReleasesMemoryAfterDurable) { WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); auto mounted = PrepareSimpleSegment(service, "batch_evict_finalize_seg"); OpLogBatchStorage storage(cluster_id, *backend); @@ -3275,6 +3366,7 @@ TEST_F(MasterServiceHATest, BatchEvictReleasesMemoryAfterDurable) { PutObjectOnSegment(service, mounted.client_id, key, "batch_evict_finalize_seg"); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); std::this_thread::sleep_for(std::chrono::milliseconds(60)); backend->BlockTxn(); @@ -3291,15 +3383,14 @@ TEST_F(MasterServiceHATest, BatchEvictReleasesMemoryAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); + EXPECT_EQ(1024, TenantUsedBytes(service)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); + EXPECT_EQ(0, TenantUsedBytes(service)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = "after_batch_evict_finalize_key"; - auto after_finalize = - service.PutStart(mounted.client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(mounted.client_id, "after_batch_evict_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } TEST_F(MasterServiceHATest, EvictDiskReplicaWritesBatchRecordOpLog) { @@ -3360,7 +3451,7 @@ TEST_F(MasterServiceHATest, EvictDiskReplicaReleasesLocalDiskAfterDurable) { .set_enable_offload(true) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); const std::string segment_name = "batch_disk_evict_finalize_segment"; auto mounted = PrepareSimpleSegment(service, segment_name); @@ -3383,6 +3474,7 @@ TEST_F(MasterServiceHATest, EvictDiskReplicaReleasesLocalDiskAfterDurable) { local_disk_replica) .has_value()); ReadBatchEventually(storage, 3, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); SetLocalDiskUsedBytesForTesting(service, mounted.client_id, 1024); backend->BlockTxn(); @@ -3402,6 +3494,8 @@ TEST_F(MasterServiceHATest, EvictDiskReplicaReleasesLocalDiskAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 4, batch); + EXPECT_EQ(1024, GetLocalDiskUsedBytesForTesting(service, segment_name)); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); EXPECT_EQ(0, GetLocalDiskUsedBytesForTesting(service, segment_name)); } @@ -3462,7 +3556,7 @@ TEST_F(MasterServiceHATest, NoFBatchEvictReleasesNoFSpaceAfterDurable) { .set_oplog_batch_max_entries(1) .build(); MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + auto* writer = InstallGatedWriter(service, backend); NoFSegment nof_segment = MakeNoFSegment("batch_nof_evict_finalize_segment", "batch_nof_evict_finalize_endpoint", @@ -3482,6 +3576,7 @@ TEST_F(MasterServiceHATest, NoFBatchEvictReleasesNoFSpaceAfterDurable) { OpLogBatchStorage storage(cluster_id, *backend); OpLogBatchRecord batch; ReadBatchEventually(storage, 1, batch); + ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); std::this_thread::sleep_for(std::chrono::milliseconds(60)); backend->BlockTxn(); @@ -3497,15 +3592,12 @@ TEST_F(MasterServiceHATest, NoFBatchEvictReleasesNoFSpaceAfterDurable) { backend->AllowTxn(); ReadBatchEventually(storage, 2, batch); + ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); - if (!before_finalize.has_value()) { - const std::string after_finalize_key = - "after_batch_nof_evict_finalize_key"; - auto after_finalize = service.PutStart(client_id, after_finalize_key, - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) - << toString(after_finalize.error()); - } + auto after_finalize = + service.PutStart(client_id, "after_batch_nof_evict_finalize_key", + kDefaultTenant, 1024, config); + EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } #endif From 8acb1e79aad6860c8534a808c2e4e3999aab5508 Mon Sep 17 00:00:00 2001 From: JieTang66 <57845979+JieTang66@users.noreply.github.com> Date: Tue, 11 Aug 2026 16:20:11 +0800 Subject: [PATCH 032/483] feat(ascend_direct): auto-detect AutoConnect & Client-Server mode via GetCapability, inject LocalCommRes (#3302) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(ascend_direct): auto-enable AutoConnect & Client-Server mode via GetCapability Unify capability resolution (AUTO_CONNECT, CLIENT_SERVER_COMM) behind ResolveAscendCapabilityFlag(): env var first, then GetCapability probe, off on older libs lacking the weak symbol. Fix AUTO_CONNECT old-build gap; add CLIENT_SERVER_COMM probe + InitParams field + docs. * test(ascend_direct): add unit tests for ResolveAscendCapabilityFlag Cover env-var precedence (1/0/invalid), GetCapability probe (supported/not-supported/failure), and both AUTO_CONNECT and CLIENT_SERVER_COMM features. Mock AdxlEngine::GetCapability with a strong symbol overriding the weak declaration. * feat(ascend_direct): auto-inject LocalCommRes for Client-Server mode When CLIENT_SERVER_COMM capability is supported and ASCEND_LOCAL_COMM_RES is not set, auto-inject LocalCommRes={"version":"1.3"} so HIXL EngineFactory selects the HixlCS engine path. User env overrides. Add tests for inject/no-inject/user-override/env-disable scenarios. * refactor(ascend_direct): simplify CLIENT_SERVER_COMM capability handling - Single-source client_server_mode: drop the redundant client_server_mode_ member and its probe in allocateLocalSegmentID (dead code — Create() routes through ParseExecutorEnvIntoInitParams which overwrites it). Resolved once in ParseExecutorEnvIntoInitParams, matching auto_connect. - Make client_server_mode purely probe-driven: no ASCEND_CLIENT_SERVER_MODE env var; enabled iff GetCapability(CLIENT_SERVER_COMM) reports support (false on older libs via the weak-symbol fallback in adxl_compat.h). - Inline ResolveAscendCapabilityFlag back into ParseExecutorEnvIntoInitParams for auto_connect — single caller after the above. Drop the helper, its declaration, and its unit tests. --------- Co-authored-by: tangjie66 --- .../transfer_executor_base.h | 1 + .../transfer_executor_base.cpp | 6 ++ .../tests/ascend_direct_transport_test.cpp | 77 ++++++++++++++++++- 3 files changed, 83 insertions(+), 1 deletion(-) diff --git a/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.h b/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.h index e07e97162e..dd3cd65611 100644 --- a/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.h +++ b/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.h @@ -60,6 +60,7 @@ class TransferExecutorBase { bool agent_mode = false; bool roce_mode = false; bool use_fabric_mem = false; + bool client_server_mode = false; }; explicit TransferExecutorBase(const InitParams& params); diff --git a/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.cpp b/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.cpp index 6f08a3257c..73dfb885f1 100644 --- a/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.cpp +++ b/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/transfer_executor_base.cpp @@ -111,6 +111,8 @@ void TransferExecutorBase::ParseExecutorEnvIntoInitParams(InitParams& params) { params.auto_connect = true; LOG(INFO) << "AutoConnect enabled by capability probe"; } + params.client_server_mode = + adxl::IsAdxlFeatureSupported(adxl::CLIENT_SERVER_COMM); char* buffer_pool = std::getenv("ASCEND_BUFFER_POOL"); if (buffer_pool && std::strcmp(buffer_pool, "0:0") != 0) { params.use_buffer_pool = true; @@ -173,6 +175,10 @@ int TransferExecutorBase::initEngines() { if (local_comm_res) { options["adxl.LocalCommRes"] = local_comm_res; LOG(INFO) << "Set LocalCommRes to:" << local_comm_res; + } else if (params_.client_server_mode) { + options["adxl.LocalCommRes"] = R"({"version":"1.3"})"; + LOG(INFO) << "Client-Server mode enabled, set LocalCommRes to " + "{\"version\":\"1.3\"}"; } options[kAutoConnect] = params_.auto_connect ? kEnabled : kDisabled; diff --git a/mooncake-transfer-engine/tests/ascend_direct_transport_test.cpp b/mooncake-transfer-engine/tests/ascend_direct_transport_test.cpp index 071b33f47b..c2588a461c 100644 --- a/mooncake-transfer-engine/tests/ascend_direct_transport_test.cpp +++ b/mooncake-transfer-engine/tests/ascend_direct_transport_test.cpp @@ -355,6 +355,9 @@ static int g_transfer_async_count = 0; static int g_register_mem_count = 0; static int g_deregister_mem_count = 0; static std::string g_last_connect_target; +static adxl::Status g_get_capability_result = adxl::SUCCESS; +static int32_t g_get_capability_value = 0; +static std::map g_last_init_options; namespace adxl_mock { void reset() { @@ -380,6 +383,15 @@ void reset() { g_register_mem_count = 0; g_deregister_mem_count = 0; g_last_connect_target.clear(); + g_get_capability_result = adxl::SUCCESS; + g_get_capability_value = 0; + g_last_init_options.clear(); +} + +void set_capability_result(adxl::Status status, int32_t value) { + std::lock_guard lock(g_mutex); + g_get_capability_result = status; + g_get_capability_value = value; } void set_connect_result(adxl::Status status) { @@ -478,6 +490,11 @@ std::string get_last_connect_target() { std::lock_guard lock(g_mutex); return g_last_connect_target; } + +std::map get_last_init_options() { + std::lock_guard lock(g_mutex); + return g_last_init_options; +} } // namespace adxl_mock } // namespace @@ -494,8 +511,12 @@ Status AdxlEngine::Initialize( const AscendString& name, const std::map& options) { (void)name; - (void)options; g_was_initialize_called = true; + g_last_init_options.clear(); + for (const auto& kv : options) { + g_last_init_options[std::string(kv.first.GetString())] = + std::string(kv.second.GetString()); + } return g_initialize_result; } @@ -611,6 +632,13 @@ Status AdxlEngine::DeregisterMem(MemHandle mem_handle) { return SUCCESS; } +Status AdxlEngine::GetCapability(FeatureType feature_type, int32_t& value) { + (void)feature_type; + std::lock_guard lock(g_mutex); + value = g_get_capability_value; + return g_get_capability_result; +} + } // namespace adxl class AscendDirectTransportTest : public ::testing::Test { @@ -2235,6 +2263,53 @@ TEST(StoreResourceConfigSplitTest, IsRoceModeEnabled_StoreRoceP2pHccs) { unsetenv("ASCEND_GLOBAL_RESOURCE_CONFIG"); } +// ----------------------------------------------------------------------------- +// Client-Server mode: when capability is supported and user did not set +// ASCEND_LOCAL_COMM_RES, Mooncake auto-injects LocalCommRes={"version":"1.3"} +// so EngineFactory selects HixlEngine (HixlCS path). +// ----------------------------------------------------------------------------- + +class ClientServerModeTest : public AscendDirectTransportTest { + protected: + void SetUp() override { + AscendDirectTransportTest::SetUp(); + unsetenv("ASCEND_LOCAL_COMM_RES"); + } + void TearDown() override { + unsetenv("ASCEND_LOCAL_COMM_RES"); + AscendDirectTransportTest::TearDown(); + } +}; + +TEST_F(ClientServerModeTest, AutoInjectLocalCommResWhenSupported) { + adxl_mock::set_capability_result(adxl::SUCCESS, 1); + auto transport = createTransport(); + ASSERT_NE(transport, nullptr); + const auto opts = adxl_mock::get_last_init_options(); + auto it = opts.find("adxl.LocalCommRes"); + ASSERT_NE(it, opts.end()); + EXPECT_EQ(it->second, R"({"version":"1.3"})"); +} + +TEST_F(ClientServerModeTest, NoInjectWhenNotSupported) { + adxl_mock::set_capability_result(adxl::SUCCESS, 0); + auto transport = createTransport(); + ASSERT_NE(transport, nullptr); + const auto opts = adxl_mock::get_last_init_options(); + EXPECT_EQ(opts.find("adxl.LocalCommRes"), opts.end()); +} + +TEST_F(ClientServerModeTest, UserEnvOverridesAutoInject) { + adxl_mock::set_capability_result(adxl::SUCCESS, 1); + setenv("ASCEND_LOCAL_COMM_RES", R"({"version":"1.2"})", 1); + auto transport = createTransport(); + ASSERT_NE(transport, nullptr); + const auto opts = adxl_mock::get_last_init_options(); + auto it = opts.find("adxl.LocalCommRes"); + ASSERT_NE(it, opts.end()); + EXPECT_EQ(it->second, R"({"version":"1.2"})"); +} + int main(int argc, char** argv) { ::testing::InitGoogleTest(&argc, argv); return RUN_ALL_TESTS(); From 5c0724d22e7f04513a3453c8b6642a5a21b80b47 Mon Sep 17 00:00:00 2001 From: Xingyuan Wu Date: Tue, 11 Aug 2026 19:50:15 +0800 Subject: [PATCH 033/483] docs: add multi-tenant deployment guide (#3378) --- .../kv-cache-sharing-and-isolation.md | 2 +- .../mooncake-store-deployment-guide.md | 75 +--------- docs/source/deployment/multi-tenancy.md | 129 ++++++++++++++++++ docs/source/design/mooncake-store.md | 25 +--- 4 files changed, 136 insertions(+), 95 deletions(-) create mode 100644 docs/source/deployment/multi-tenancy.md diff --git a/docs/source/deployment/kv-cache-sharing-and-isolation.md b/docs/source/deployment/kv-cache-sharing-and-isolation.md index cb7706f166..79d172a6c3 100644 --- a/docs/source/deployment/kv-cache-sharing-and-isolation.md +++ b/docs/source/deployment/kv-cache-sharing-and-isolation.md @@ -123,7 +123,7 @@ Use a new namespace whenever cache compatibility may have changed, including cha Mooncake tenant configuration is independent of the framework-level model, release, and request namespaces. When the master is started with `--enable_multi_tenants=true`, the client `tenant_id` selects a tenant-scoped object namespace and the master applies that tenant's quota during admission. -Use the same `tenant_id` for framework instances that should share one quota and tenant namespace. See [Tenant Quota Management](mooncake-store-deployment-guide.md#tenant-quota-management) for configuration details. +Use the same `tenant_id` for framework instances that should share one quota and tenant namespace. See [Multi-Tenant Deployment](multi-tenancy) for configuration details. ## Operational Checklist diff --git a/docs/source/deployment/mooncake-store-deployment-guide.md b/docs/source/deployment/mooncake-store-deployment-guide.md index d4d055ac58..7ac43272dc 100644 --- a/docs/source/deployment/mooncake-store-deployment-guide.md +++ b/docs/source/deployment/mooncake-store-deployment-guide.md @@ -402,78 +402,11 @@ When tenant quota is enabled, `/metrics` also includes per-tenant quota gauges a ## Tenant Quota Management -Tenant quota admission is disabled by default. Enable strict multi-tenant mode on the master when you want memory writes admitted against connector-managed per-tenant quota: - -```bash -mooncake_master \ - --enable_multi_tenants=true \ - --tenant_quota_connector_type=file \ - --tenant_quota_connector_uri=/etc/mooncake/tenant_quotas.yaml -``` - -You can also store the same YAML policy in etcd when Mooncake Store is built with `STORE_USE_ETCD=ON`: - -```bash -mooncake_master \ - --enable_multi_tenants=true \ - --cluster_id=mooncake_cluster \ - --tenant_quota_connector_type=etcd \ - --tenant_quota_connector_uri=127.0.0.1:2379 -``` - -The etcd connector stores the policy at `mooncake-store//tenant_quota_policy`. If the key does not exist, the master starts with an empty policy so the first tenant policy can be created through the admin API. It shares the process-wide store etcd client used by HA/oplog, so if HA or oplog also uses etcd, `tenant_quota_connector_uri` must match those etcd endpoints. The policy must use schema version `1`; tenant names must be non-empty, unique, must not start with `_`, and must not contain NUL or control characters; quotas must be positive integers with optional `B`, `KB`, `MB`, `GB`, or `TB` units: - -```yaml -version: 1 - -tenants: - - name: tenant-a - quota: 200GB - - - name: tenant-b - quota: 500GB -``` - -When strict multi-tenant mode is enabled, write requests must include a registered tenant. The `default` tenant is not special unless it is explicitly registered in the connector policy. - -The same HTTP port used for metrics exposes the tenant quota admin API: - -```bash -# List tenant quota snapshots -curl -s http://:9003/api/v1/tenant_quotas - -# Query one tenant -curl -s "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" - -# Upsert an explicit policy. Explicit tenant policies must be positive. -curl -s -X PUT "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" \ - -H 'Content-Type: application/json' \ - -d '{"requested_quota_bytes":2147483648}' - -# Delete an explicit policy. The tenant must not own objects or quota usage. -curl -s -X DELETE "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" -``` - -Each tenant quota snapshot returns: - -```json -{ - "success": true, - "data": { - "tenant_id": "tenant-a", - "requested_quota_bytes": 2147483648, - "effective_quota_bytes": 2147483648, - "used_bytes": 0, - "reserved_bytes": 0, - "committed_count": 0, - "metadata_object_count": 0, - "over_quota": false, - "has_explicit_policy": true - } -} -``` +:::{toctree} +:maxdepth: 1 -In HA mode, quota admin requests are served only by the active master service. Standby, candidate, or inactive services return HTTP 503. If strict multi-tenant mode is disabled, the quota admin API returns HTTP 409 with `UNAVAILABLE_IN_CURRENT_MODE`. Deleting a non-empty tenant returns HTTP 409 with `TENANT_NOT_EMPTY`. +Multi-Tenant Deployment +::: --- diff --git a/docs/source/deployment/multi-tenancy.md b/docs/source/deployment/multi-tenancy.md new file mode 100644 index 0000000000..b4c48cf306 --- /dev/null +++ b/docs/source/deployment/multi-tenancy.md @@ -0,0 +1,129 @@ +# Multi-Tenant Deployment + +## Configure the Master + +### File Connector + +Tenant quota admission is disabled by default. Enable strict multi-tenant mode on the master when you want memory writes admitted against connector-managed per-tenant quota: + +```bash +mooncake_master \ + --enable_multi_tenants=true \ + --tenant_quota_connector_type=file \ + --tenant_quota_connector_uri=/etc/mooncake/tenant_quotas.yaml +``` + +### etcd Connector + +You can also store the same YAML policy in etcd when Mooncake Store is built with `STORE_USE_ETCD=ON`: + +```bash +mooncake_master \ + --enable_multi_tenants=true \ + --cluster_id=mooncake_cluster \ + --tenant_quota_connector_type=etcd \ + --tenant_quota_connector_uri=127.0.0.1:2379 +``` + +The etcd connector stores the policy at `mooncake-store//tenant_quota_policy`. If the key does not exist, the master starts with an empty policy so the first tenant policy can be created through the admin API. It shares the process-wide store etcd client used by HA/oplog, so if HA or oplog also uses etcd, `tenant_quota_connector_uri` must match those etcd endpoints. + +## Define the Tenant Policy + +The policy must use schema version `1`; tenant names must be non-empty, unique, must not start with `_`, and must not contain NUL or control characters; quotas must be positive integers with optional `B`, `KB`, `MB`, `GB`, or `TB` units: + +```yaml +version: 1 + +tenants: + - name: tenant-a + quota: 200GB + + - name: tenant-b + quota: 500GB +``` + +When strict multi-tenant mode is enabled, write requests must include a registered tenant. The `default` tenant is not special unless it is explicitly registered in the connector policy. + +## Manage Tenant Quotas + +The same HTTP port used for metrics exposes the tenant quota admin API: + +```bash +# List tenant quota snapshots +curl -s http://:9003/api/v1/tenant_quotas + +# Query one tenant +curl -s "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" + +# Upsert an explicit policy. Explicit tenant policies must be positive. +curl -s -X PUT "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" \ + -H 'Content-Type: application/json' \ + -d '{"requested_quota_bytes":2147483648}' + +# Delete an explicit policy. The tenant must not own objects or quota usage. +curl -s -X DELETE "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" +``` + +Each tenant quota snapshot returns: + +```json +{ + "success": true, + "data": { + "tenant_id": "tenant-a", + "requested_quota_bytes": 2147483648, + "effective_quota_bytes": 2147483648, + "used_bytes": 0, + "reserved_bytes": 0, + "committed_count": 0, + "metadata_object_count": 0, + "over_quota": false, + "has_explicit_policy": true + } +} +``` + +In HA mode, quota admin requests are served only by the active master service. Standby, candidate, or inactive services return HTTP 503. If strict multi-tenant mode is disabled, the quota admin API returns HTTP 409 with `UNAVAILABLE_IN_CURRENT_MODE`. Deleting a non-empty tenant returns HTTP 409 with `TENANT_NOT_EMPTY`. + +## SGLang + +When Mooncake is used as the HiCache storage backend, set `tenant_id` in the +Mooncake backend configuration: + +```bash +--hicache-storage-backend mooncake \ +--hicache-storage-backend-extra-config \ + '{"master_server_address":"127.0.0.1:50051","tenant_id":"tenant-a"}' +``` + +Alternatively, add `tenant_id` to the JSON file selected by +`SGLANG_HICACHE_MOONCAKE_CONFIG_PATH`, or use `MOONCAKE_TENANT_ID` when loading +the Mooncake configuration from environment variables. SGLang forwards the +resolved value to the Mooncake client. + +All prefill, decode, and replica instances that should share KV cache entries +must use the same `tenant_id` and compatible model and release namespaces. + +## vLLM + +Add `tenant_id` to the Mooncake client JSON configuration: + +```json +"tenant_id": "tenant-a" +``` + +Point `MOONCAKE_CONFIG_PATH` at that file and enable +`MooncakeStoreConnector` through `--kv-transfer-config`: + +```bash +MOONCAKE_CONFIG_PATH=/path/to/mooncake_config.json \ +vllm serve \ + --kv-transfer-config \ + '{"kv_connector":"MooncakeStoreConnector","kv_role":"kv_both"}' +``` + +`MooncakeStoreConnector` reads the JSON during initialization and passes its +tenant ID to the Mooncake client. + +All prefill, decode, and replica instances that should share KV cache entries +must use the same `tenant_id` and compatible model and release namespaces. diff --git a/docs/source/design/mooncake-store.md b/docs/source/design/mooncake-store.md index dee3f9a651..557d40ed88 100644 --- a/docs/source/design/mooncake-store.md +++ b/docs/source/design/mooncake-store.md @@ -95,19 +95,7 @@ To reduce cache warm-up time after a master restart, the Master Service supports The Master Service can optionally enforce strict multi-tenant memory quota admission. This feature is disabled by default. When `enable_multi_tenants=false`, request tenant IDs are ignored for object placement, all objects use the `default` namespace, and tenant quota management requests return `UNAVAILABLE_IN_CURRENT_MODE`. -When strict multi-tenant mode is enabled, the tenant quota policy is loaded from the configured connector. Supported connector types are `file` and, when the store is built with `STORE_USE_ETCD=ON`, `etcd`. The `file` connector uses `tenant_quota_connector_uri=` as a writable YAML policy path. The `etcd` connector uses `tenant_quota_connector_uri=` as the etcd endpoints string and stores the same YAML policy in `mooncake-store//tenant_quota_policy`; if that key does not exist, the master starts with an empty policy so the first policy can be created through the admin API. The etcd connector shares the process-wide store etcd client used by HA/oplog, so deployments that enable both must configure matching etcd endpoints. Tenants must be explicitly present in that connector policy before they can write. Missing tenants, empty tenants, and an unregistered `default` tenant are rejected with `TENANT_NOT_REGISTERED`. - -The YAML policy uses schema version `1`: - -```yaml -version: 1 - -tenants: - - name: tenant-a - quota: 200GB -``` - -Tenant names must be non-empty, unique, must not start with `_`, and must not contain NUL or control characters. Quotas must be positive integers and may use `B`, `KB`, `MB`, `GB`, or `TB` units. +See [Multi-Tenant Deployment](../deployment/multi-tenancy.md) for configuration details. Effective quota is recomputed from the current registered memory capacity: @@ -118,16 +106,7 @@ Effective quota is recomputed from the current registered memory capacity: `PutStart` and size-changing `UpsertStart` charge quota before memory is allocated. If the first reservation fails, the master performs tenant-scoped memory eviction for the target tenant and retries the reservation. The retry is bounded to two eviction attempts. Tenant quota eviction scans only the target tenant, skips hard-pinned objects, honors soft-pin eviction configuration, and preserves grouped-object lease safety checks. -Admin policy changes are persisted before the final in-memory policy is applied. `PUT` writes the connector first and then applies the policy in memory. `DELETE` first marks the tenant unregistered in memory to block concurrent writes, verifies the tenant is empty, writes the connector, and rolls back the in-memory mark if the connector write fails. The admin HTTP API exposes: - -| Method | Path | Description | -|--------|------|-------------| -| `GET` | `/api/v1/tenant_quotas` | List quota snapshots for active or explicit tenants | -| `GET` | `/api/v1/tenant_quotas?tenant_id=` | Query one tenant quota snapshot | -| `PUT` | `/api/v1/tenant_quotas?tenant_id=` | Create or update a tenant quota policy | -| `DELETE` | `/api/v1/tenant_quotas?tenant_id=` | Delete an empty tenant quota policy | - -Tenant quota snapshots include `tenant_id`, `requested_quota_bytes`, `effective_quota_bytes`, `used_bytes`, `reserved_bytes`, `committed_count`, `metadata_object_count`, `over_quota`, and `has_explicit_policy`. +Admin policy changes are persisted before the final in-memory policy is applied. `PUT` writes the connector first and then applies the policy in memory. `DELETE` first marks the tenant unregistered in memory to block concurrent writes, verifies the tenant is empty, writes the connector, and rolls back the in-memory mark if the connector write fails. Snapshots restore object runtime state only. Tenant quota policy is always loaded from the connector after metadata restore, then usage and effective quota are rebuilt from restored metadata and current registered capacity. If the connector cannot be loaded in strict multi-tenant mode, startup fails. From 1b6b711850c6501156b56504be2445b1c6cafc0e Mon Sep 17 00:00:00 2001 From: Stary Date: Tue, 11 Aug 2026 22:48:10 +0800 Subject: [PATCH 034/483] [Doc] Align pre-commit guidance to PR-changed files (#3277) * [Doc] Align pre-commit guidance to PR-changed files Recommend running pre-commit on staged or PR-scoped files instead of --all-files, so routine PRs do not pick up unrelated historical rewrites. Signed-off-by: staryxchen Co-authored-by: Cursor * [Doc] Fix PR-scoped pre-commit base to origin/main Use `origin/main...HEAD` instead of local `main...HEAD` so the diff base always reflects the upstream tip; prepend `git fetch origin main` to keep the remote ref fresh before resolving the file list. Local `main` lags behind when working long on a feature branch, which would include upstream files unrelated to the PR. Addresses Aionw's review comment on PR#3277. Signed-off-by: staryxchen * [Doc] Fix PR-scoped pre-commit base in CONTRIBUTING.md Sync the PR-scoped pre-commit command in CONTRIBUTING.md with the preceding fix in .pre-commit-config.yaml: use `origin/main...HEAD` instead of local `main...HEAD` and prepend `git fetch origin main` so the remote ref stays fresh before resolving the file list. Part of PR#3277 review feedback from Aionw. Signed-off-by: staryxchen --------- Signed-off-by: staryxchen Co-authored-by: Cursor --- .github/pull_request_template.md | 2 +- .pre-commit-config.yaml | 4 +++- AGENTS.md | 8 +++++--- CONTRIBUTING.md | 18 +++++++++++++++--- 4 files changed, 24 insertions(+), 8 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index a926d40787..2e44ed37cf 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -45,7 +45,7 @@ - [ ] I have performed a self-review of my own code - [ ] I have formatted my code using `./scripts/code_format.sh` -- [ ] I have run `pre-commit run --all-files` and all hooks pass +- [ ] I have run pre-commit on the files changed in this PR and all hooks pass - [ ] I have updated the documentation (if applicable) - [ ] I have added tests to prove my changes are effective - [ ] For changes >500 LOC: I have filed an RFC issue diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index e29bb9c9be..46607d9270 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,8 @@ # Pre-commit hooks configuration for Mooncake # Install: pip install -r requirements-dev.txt && pre-commit install -# Run manually on all files: pre-commit run --all-files +# Staged files: pre-commit run +# PR-changed files: git fetch origin main && pre-commit run --files $(git diff --name-only --diff-filter=ACMR origin/main...HEAD) +# Full-repo (intentional cleanup only): pre-commit run --all-files # Format all C/C++ files explicitly: ./scripts/code_format.sh --all # Note: clang-format should already be available (installed via system packages or dependencies.sh) # Exclusions: build artifacts, vendored extern code, generated wheels. diff --git a/AGENTS.md b/AGENTS.md index f7332c2d0a..6ab75b16b9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -18,8 +18,10 @@ checklist, and AI assistance disclosure. - For AI-assisted changes, make sure the human submitter has reviewed every changed line and can defend the change end-to-end. -- Run pre-commit locally on the files touched by the change before handoff when - the toolchain is available. If broader hooks or `pre-commit run --all-files` - rewrite unrelated files, do not include those unrelated edits in the PR. +- Before handoff, run pre-commit on the files touched by the change when the + toolchain is available (see `CONTRIBUTING.md` for the PR-scoped + `pre-commit run --files ...` command). Do not use + `pre-commit run --all-files` for routine PRs; if it rewrites unrelated + files, leave those edits out of the PR. - Keep PRs lean: review `git diff` before staging, and include only changes required for the requested task. diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 68d0741220..fae47e6fed 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -65,9 +65,21 @@ committing again. Use `./scripts/code_format.sh --all` only when intentionally formatting the whole project. #### Usage -Run hooks on all files (the first run installs hook environments). The C/C++ -hook remains limited to staged line ranges; use `./scripts/code_format.sh --all` -for an intentional whole-project C/C++ format: +After `pre-commit install`, hooks run on each commit. To run them manually on +staged files (the first run may install hook environments): +```bash +pre-commit run +``` +Before opening a PR, run hooks only on files changed against the PR base +(default `origin/main`). The C/C++ hook remains limited to staged or changed +line ranges: +```bash +git fetch origin main +pre-commit run --files $(git diff --name-only --diff-filter=ACMR origin/main...HEAD) +``` +Use full-repo checks only for intentional whole-project cleanup; do not fold +unrelated rewrites into a feature PR. Prefer `./scripts/code_format.sh --all` +for a deliberate whole-project C/C++ format: ```bash pre-commit run --all-files ``` From da2f5be07f1901ba1825da0caa94d65698dd2726 Mon Sep 17 00:00:00 2001 From: SongOf <46475785+SongOf@users.noreply.github.com> Date: Tue, 11 Aug 2026 22:51:18 +0800 Subject: [PATCH 035/483] [Bugfix][TE] Reject MC_IB_PORT=0 instead of disabling every RNIC (#3381) Co-authored-by: maxlisongsong --- mooncake-transfer-engine/src/config.cpp | 4 +- .../tests/config_test.cpp | 50 +++++++++++++++++++ 2 files changed, 53 insertions(+), 1 deletion(-) diff --git a/mooncake-transfer-engine/src/config.cpp b/mooncake-transfer-engine/src/config.cpp index e3b0154659..64a08332e1 100644 --- a/mooncake-transfer-engine/src/config.cpp +++ b/mooncake-transfer-engine/src/config.cpp @@ -128,7 +128,9 @@ void loadGlobalConfig(GlobalConfig& config) { const char* port_env = std::getenv("MC_IB_PORT"); if (port_env) { int val = atoi(port_env); - if (val >= 0 && val < 256) + // IB port numbers are 1-based. Accepting 0 made ibv_query_port fail on + // every device, disabling the whole topology. + if (val > 0 && val < 256) config.port = uint8_t(val); else LOG(WARNING) << "Ignore value from environment variable MC_IB_PORT"; diff --git a/mooncake-transfer-engine/tests/config_test.cpp b/mooncake-transfer-engine/tests/config_test.cpp index e3f75d5c6b..f81a65d39a 100644 --- a/mooncake-transfer-engine/tests/config_test.cpp +++ b/mooncake-transfer-engine/tests/config_test.cpp @@ -277,6 +277,56 @@ TEST_F(ConnPauseTtlEnvTest, EmptyStringKeepsDefault) { EXPECT_EQ(config.conn_pause_ttl_ms, 17); } +// MC_IB_PORT names the RDMA port opened on every device. Port numbers are +// 1-based, so 0 is not a "disable" value: it makes ibv_query_port fail on +// every device and takes the whole topology down. +class IbPortEnvTest : public ::testing::Test { + protected: + void TearDown() override { ::unsetenv("MC_IB_PORT"); } +}; + +TEST_F(IbPortEnvTest, DefaultIsPortOneWhenUnset) { + ::unsetenv("MC_IB_PORT"); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 1); +} + +TEST_F(IbPortEnvTest, ValidOverrideIsApplied) { + ASSERT_EQ(::setenv("MC_IB_PORT", "2", 1), 0); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 2); +} + +TEST_F(IbPortEnvTest, ZeroIsRejected) { + ASSERT_EQ(::setenv("MC_IB_PORT", "0", 1), 0); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 1); +} + +TEST_F(IbPortEnvTest, OutOfRangeIsRejected) { + ASSERT_EQ(::setenv("MC_IB_PORT", "256", 1), 0); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 1); +} + +TEST_F(IbPortEnvTest, NegativeIsRejected) { + ASSERT_EQ(::setenv("MC_IB_PORT", "-1", 1), 0); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 1); +} + +TEST_F(IbPortEnvTest, NonNumericIsRejected) { + ASSERT_EQ(::setenv("MC_IB_PORT", "abc", 1), 0); + GlobalConfig config; + loadGlobalConfig(config); + EXPECT_EQ(config.port, 1); +} + // MC_MAX_CONCURRENT_REG_MR caps how many buffers registerLocalMemoryBatch() // registers at once; 0 (the default) means unbounded. 0 is therefore also what // a silent atol() fallback would produce on a typo, which would read as "the From e05d4290829b67bd8e3070c312c3ff9b143a62e7 Mon Sep 17 00:00:00 2001 From: mokeke <103873176+mo-ke-ke@users.noreply.github.com> Date: Wed, 12 Aug 2026 11:51:45 +0800 Subject: [PATCH 036/483] [Store] Fix current-stream readiness for DummyClient CUDA IPC tensor writes (#3303) Co-authored-by: mo-ke-ke --- .../store/store_py_parallel_write.h | 95 +++++- .../include/device/cuda_ipc_buffer.h | 3 + mooncake-store/src/device/cuda_ipc_buffer.cpp | 31 ++ mooncake-wheel/tests/test_dummy_client.py | 288 +++++++++++++++++- 4 files changed, 406 insertions(+), 11 deletions(-) diff --git a/mooncake-integration/store/store_py_parallel_write.h b/mooncake-integration/store/store_py_parallel_write.h index 9f23e3d501..b31e698918 100644 --- a/mooncake-integration/store/store_py_parallel_write.h +++ b/mooncake-integration/store/store_py_parallel_write.h @@ -19,16 +19,38 @@ std::optional> try_dummy_cuda_ipc_batch_put_tensor_impl( const std::vector &infos, const ReplicateConfig &config) { if (!use_dummy_client_ || keys.size() != infos.size()) return std::nullopt; + struct CudaStreamContext { + int32_t device_id; + uintptr_t stream_handle; + }; + std::vector> cuda_stream_contexts( + infos.size()); + for (size_t i = 0; i < infos.size(); ++i) { + const py::object &owner = infos[i].owner; + if (!infos[i].valid() || !owner || owner.is_none()) continue; + + try { + if (!owner.attr("is_cuda").cast()) continue; + cuda_stream_contexts[i] = CudaStreamContext{ + .device_id = owner.attr("get_device")().cast(), + .stream_handle = + torch_module() + .attr("cuda") + .attr("current_stream")(owner.attr("device")) + .attr("cuda_stream") + .cast(), + }; + } catch (const py::error_already_set &) { + PyErr_Clear(); + } catch (const py::cast_error &) { + PyErr_Clear(); + } + } + std::vector results(keys.size(), 0); py::gil_scoped_release release_gil; - std::vector write_requests; - std::vector original_indices; - std::vector> metadata_allocations; - write_requests.reserve(infos.size()); - original_indices.reserve(infos.size()); - metadata_allocations.reserve(infos.size()); - + std::vector> ipc_payloads(infos.size()); for (size_t i = 0; i < infos.size(); ++i) { if (!infos[i].valid()) { results[i] = to_py_ret(ErrorCode::INVALID_PARAMS); @@ -40,6 +62,39 @@ std::optional> try_dummy_cuda_ipc_batch_put_tensor_impl( reinterpret_cast(infos[i].data_ptr), infos[i].tensor_size); if (!payload) return std::nullopt; + ipc_payloads[i] = std::move(*payload); + } + + bool stream_context_invalid = false; + for (size_t i = 0; i < infos.size(); ++i) { + if (!ipc_payloads[i].has_value()) continue; + if (!cuda_stream_contexts[i].has_value() || + cuda_stream_contexts[i]->device_id != ipc_payloads[i]->device_id) { + results[i] = to_py_ret(ErrorCode::INTERNAL_ERROR); + stream_context_invalid = true; + } + } + if (stream_context_invalid) { + LOG(ERROR) << "CUDA IPC tensor stream context validation failed"; + for (size_t i = 0; i < infos.size(); ++i) { + if (ipc_payloads[i].has_value()) { + results[i] = to_py_ret(ErrorCode::INTERNAL_ERROR); + } + } + return results; + } + + std::vector write_requests; + std::vector original_indices; + std::vector> metadata_allocations; + std::vector> unique_streams; + write_requests.reserve(infos.size()); + original_indices.reserve(infos.size()); + metadata_allocations.reserve(infos.size()); + unique_streams.reserve(infos.size()); + + for (size_t i = 0; i < infos.size(); ++i) { + if (!ipc_payloads[i].has_value()) continue; size_t metadata_size = infos[i].metadata.header.data_offset; auto metadata = store_->allocate_client_buffer(metadata_size); @@ -56,17 +111,41 @@ std::optional> try_dummy_cuda_ipc_batch_put_tensor_impl( .ptr = reinterpret_cast(metadata->ptr()), .size = static_cast(metadata_size), }, - .payload = *payload, + .payload = *ipc_payloads[i], }); original_indices.push_back(i); metadata_allocations.push_back( std::make_unique(std::move(*metadata))); + + const CudaStreamContext &stream_context = *cuda_stream_contexts[i]; + bool stream_is_unique = true; + for (const auto &[device_id, stream_handle] : unique_streams) { + if (device_id == stream_context.device_id && + stream_handle == stream_context.stream_handle) { + stream_is_unique = false; + break; + } + } + if (stream_is_unique) { + unique_streams.emplace_back(stream_context.device_id, + stream_context.stream_handle); + } } if (!write_requests.empty()) { ReplicateConfig write_config = MakeIndexedConfig(config, original_indices); auto dummy_client = std::static_pointer_cast(store_); + for (const auto &[device_id, stream_handle] : unique_streams) { + if (!mooncake::device::SynchronizeCudaStream(device_id, + stream_handle)) { + LOG(ERROR) << "CUDA IPC tensor stream synchronization failed"; + for (size_t index : original_indices) { + results[index] = to_py_ret(ErrorCode::INTERNAL_ERROR); + } + return results; + } + } std::vector op_results = dummy_client->batch_put_from_cuda_ipc(write_requests, write_config); if (!apply_indexed_results("put", op_results, original_indices, diff --git a/mooncake-store/include/device/cuda_ipc_buffer.h b/mooncake-store/include/device/cuda_ipc_buffer.h index 3a6821482f..4a53c7fa16 100644 --- a/mooncake-store/include/device/cuda_ipc_buffer.h +++ b/mooncake-store/include/device/cuda_ipc_buffer.h @@ -11,6 +11,9 @@ namespace device { tl::expected ExportCudaIpcBuffer( const void *ptr, size_t size); +tl::expected SynchronizeCudaStream(int32_t device_id, + uintptr_t stream_handle); + class CudaIpcBufferMapping { public: CudaIpcBufferMapping() = default; diff --git a/mooncake-store/src/device/cuda_ipc_buffer.cpp b/mooncake-store/src/device/cuda_ipc_buffer.cpp index 5953ff57d4..317dcc040e 100644 --- a/mooncake-store/src/device/cuda_ipc_buffer.cpp +++ b/mooncake-store/src/device/cuda_ipc_buffer.cpp @@ -122,6 +122,37 @@ tl::expected ExportCudaIpcBuffer( #endif } +tl::expected SynchronizeCudaStream(int32_t device_id, + uintptr_t stream_handle) { +#if defined(USE_CUDA) + int current_device = -1; + if (cudaGetDevice(¤t_device) != cudaSuccess) { + ClearCudaError(); + return tl::unexpected(ErrorCode::INTERNAL_ERROR); + } + + cudaError_t operation_status = cudaSetDevice(device_id); + if (operation_status == cudaSuccess) { + const cudaStream_t stream = + stream_handle == 0 ? nullptr + : reinterpret_cast(stream_handle); + operation_status = cudaStreamSynchronize(stream); + } + if (operation_status != cudaSuccess) ClearCudaError(); + + const cudaError_t restore_status = cudaSetDevice(current_device); + if (restore_status != cudaSuccess) ClearCudaError(); + if (operation_status != cudaSuccess || restore_status != cudaSuccess) { + return tl::unexpected(ErrorCode::INTERNAL_ERROR); + } + return {}; +#else + (void)device_id; + (void)stream_handle; + return tl::unexpected(ErrorCode::INVALID_PARAMS); +#endif +} + CudaIpcBufferMapping::~CudaIpcBufferMapping() { Close(); } CudaIpcBufferMapping::CudaIpcBufferMapping( diff --git a/mooncake-wheel/tests/test_dummy_client.py b/mooncake-wheel/tests/test_dummy_client.py index bbc0adddaf..dc7138b448 100644 --- a/mooncake-wheel/tests/test_dummy_client.py +++ b/mooncake-wheel/tests/test_dummy_client.py @@ -1,11 +1,15 @@ -import unittest +import json +import math import os -import time import threading +import time +import unittest try: import torch -except ImportError: +except ModuleNotFoundError as error: + if error.name != "torch": + raise torch = None from mooncake.store import MooncakeDistributedStore, SoftPinAction @@ -16,6 +20,262 @@ # Use environment variable if set, otherwise use default default_kv_lease_ttl = int(os.getenv("DEFAULT_KV_LEASE_TTL", DEFAULT_DEFAULT_KV_LEASE_TTL)) +CUDA_IPC_STREAM_READINESS_PAYLOAD_BYTES = 16 * 1024 * 1024 +_CUDA_IPC_INITIAL_BYTE = 0x31 +_CUDA_IPC_FINAL_BYTE = 0xA7 +_CUDA_LAUNCH_BLOCKING_DISABLED_VALUES = {"", "0", "false", "no", "off"} + + +class _CudaIpcStreamReadinessFailure(AssertionError): + """A CUDA readiness failure whose diagnostics are safe to publish.""" + + def __init__(self, diagnostics): + self.diagnostics = diagnostics + compact = json.dumps(diagnostics, sort_keys=True, separators=(",", ":")) + super().__init__(f"CUDA IPC stream readiness case failed: {compact}") + + +def _cuda_launch_blocking_is_active(): + value = os.environ.get("CUDA_LAUNCH_BLOCKING") + return ( + value is not None + and value.strip().lower() not in _CUDA_LAUNCH_BLOCKING_DISABLED_VALUES + ) + + +def _cuda_stream_readiness_skip_reason(): + if torch is None: + return "PyTorch is not available" + if getattr(torch.version, "cuda", None) is None: + return "PyTorch does not have CUDA support" + if not torch.cuda.is_available(): + return "CUDA is not available" + if not callable(getattr(torch.cuda, "_sleep", None)): + return "torch.cuda._sleep is not available" + if _cuda_launch_blocking_is_active(): + return "CUDA_LAUNCH_BLOCKING disables the asynchronous window" + return None + + +def _calibrate_cuda_sleep_cycles(): + target_ms = 300.0 + minimum_ms = 200.0 + maximum_ms = 750.0 + minimum_cycles = 1 + maximum_cycles = 2_000_000_000 + cycles = 1_000_000 + device = torch.device("cuda", torch.cuda.current_device()) + calibration_stream = torch.cuda.Stream(device=device) + + for _ in range(6): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + with torch.cuda.stream(calibration_stream): + start.record() + torch.cuda._sleep(cycles) + end.record() + end.synchronize() + elapsed_ms = float(start.elapsed_time(end)) + if minimum_ms <= elapsed_ms <= maximum_ms: + return cycles + if not math.isfinite(elapsed_ms) or elapsed_ms <= 0.0: + break + cycles = round(cycles * target_ms / elapsed_ms) + cycles = max(minimum_cycles, min(maximum_cycles, cycles)) + + raise AssertionError("Unable to establish the CUDA asynchronous window") + + +def _put_cuda_readiness_tensors(store, keys, tensors, batch_width): + if batch_width == 1: + return store.put_tensor(keys[0], tensors[0]) == 0 + return list(store.batch_put_tensor(keys, tensors)) == [0] * batch_width + + +def _get_cuda_readiness_tensors(store, keys, batch_width): + if batch_width == 1: + return [store.get_tensor(keys[0])] + return list(store.batch_get_tensor(keys)) + + +def _validate_cuda_readiness_payloads( + retrieved, batch_width, payload_bytes, expected_byte +): + if len(retrieved) != batch_width: + return False, -1, 0 + + metadata_ok = True + mismatch_count = 0 + checksum = 0 + for actual in retrieved: + if not isinstance(actual, torch.Tensor): + metadata_ok = False + mismatch_count = -1 + continue + actual_cpu = actual.detach().cpu() + checksum += int(actual_cpu.sum(dtype=torch.int64).item()) + if actual_cpu.dtype != torch.uint8 or tuple(actual_cpu.shape) != ( + payload_bytes, + ): + metadata_ok = False + mismatch_count = -1 + continue + if mismatch_count >= 0: + mismatch_count += int( + torch.count_nonzero(actual_cpu != expected_byte).item() + ) + return metadata_ok, mismatch_count, checksum + + +def _new_cuda_readiness_diagnostics(batch_width, explicit_sync): + return { + "mode": "control" if explicit_sync else "async", + "batch_width": batch_width, + "pending_at_call": None, + "ready_at_return": None, + "latency_ms": None, + "mismatch_count": None, + "checksum": None, + } + + +def _run_dummy_cuda_ipc_stream_readiness_case( + store, + batch_width, + explicit_sync, + sleep_cycles, + payload_bytes=CUDA_IPC_STREAM_READINESS_PAYLOAD_BYTES, +): + if batch_width not in (1, 2) or payload_bytes <= 0: + raise ValueError("invalid synthetic CUDA readiness parameters") + + diagnostics = _new_cuda_readiness_diagnostics(batch_width, explicit_sync) + prefix = f"dummy_cuda_ipc_readiness_{os.getpid()}_{time.monotonic_ns()}" + warm_keys = [f"{prefix}_warm_{index}" for index in range(batch_width)] + measured_keys = [f"{prefix}_measured_{index}" for index in range(batch_width)] + cleanup_keys = [*warm_keys, *measured_keys] + tensors = [] + primary_failed = False + cleanup_failed = False + + try: + device = torch.device("cuda", torch.cuda.current_device()) + with torch.cuda.device(device): + current = torch.cuda.current_stream(device) + tensors = [ + torch.full( + (payload_bytes,), + _CUDA_IPC_INITIAL_BYTE, + dtype=torch.uint8, + device=device, + ).contiguous() + for _ in range(batch_width) + ] + + # Make the initial value host-ready before opening the async window. + current.synchronize() + if not all( + tensor.is_contiguous() + and tensor.dtype == torch.uint8 + and tuple(tensor.shape) == (payload_bytes,) + for tensor in tensors + ): + primary_failed = True + + if not primary_failed: + warm_ok = _put_cuda_readiness_tensors( + store, warm_keys, tensors, batch_width + ) + warm_payloads = _get_cuda_readiness_tensors( + store, warm_keys, batch_width + ) + warm_metadata_ok, warm_mismatch_count, _ = ( + _validate_cuda_readiness_payloads( + warm_payloads, + batch_width, + payload_bytes, + _CUDA_IPC_INITIAL_BYTE, + ) + ) + if not warm_ok or not warm_metadata_ok or warm_mismatch_count != 0: + primary_failed = True + + if not primary_failed: + producer = torch.cuda.Stream(device=device) + ready = torch.cuda.Event() + try: + with torch.cuda.stream(producer): + torch.cuda._sleep(sleep_cycles) + for tensor in tensors: + tensor.fill_(_CUDA_IPC_FINAL_BYTE) + ready.record(producer) + + current.wait_stream(producer) + if explicit_sync: + current.synchronize() + + ready_before_put = bool(ready.query()) + diagnostics["pending_at_call"] = not ready_before_put + valid_window = ( + ready_before_put if explicit_sync else not ready_before_put + ) + if not valid_window: + primary_failed = True + else: + put_started = time.perf_counter() + try: + put_ok = _put_cuda_readiness_tensors( + store, measured_keys, tensors, batch_width + ) + finally: + diagnostics["latency_ms"] = round( + (time.perf_counter() - put_started) * 1000.0, 3 + ) + diagnostics["ready_at_return"] = bool(ready.query()) + if not put_ok or not diagnostics["ready_at_return"]: + primary_failed = True + finally: + producer.synchronize() + + if not primary_failed: + retrieved = _get_cuda_readiness_tensors( + store, measured_keys, batch_width + ) + metadata_ok, mismatch_count, checksum = ( + _validate_cuda_readiness_payloads( + retrieved, + batch_width, + payload_bytes, + _CUDA_IPC_FINAL_BYTE, + ) + ) + diagnostics["mismatch_count"] = mismatch_count + diagnostics["checksum"] = checksum + expected_checksum = _CUDA_IPC_FINAL_BYTE * payload_bytes * batch_width + if ( + not metadata_ok + or mismatch_count != 0 + or checksum != expected_checksum + ): + primary_failed = True + except Exception: + primary_failed = True + finally: + for key in cleanup_keys: + try: + if store.remove(key, force=True) != 0: + cleanup_failed = True + except Exception: + cleanup_failed = True + + # Preserve primary failure precedence while suppressing raw exception text, + # which may contain identifiers forbidden from public diagnostics. + if primary_failed: + raise _CudaIpcStreamReadinessFailure(diagnostics) from None + if cleanup_failed: + raise _CudaIpcStreamReadinessFailure(diagnostics) from None + return diagnostics + def get_client(store, local_buffer_size_param=None): """Initialize and setup the distributed store client.""" @@ -483,6 +743,28 @@ def test_tensor_operations(self): for cleanup_key in cleanup_keys: self.store.remove(cleanup_key) + def _run_dummy_cuda_ipc_stream_readiness_regression(self, batch_width): + skip_reason = _cuda_stream_readiness_skip_reason() + if skip_reason is not None: + self.skipTest(skip_reason) + + sleep_cycles = _calibrate_cuda_sleep_cycles() + for explicit_sync in (False, True): + mode = "control" if explicit_sync else "async" + with self.subTest(mode=mode, batch_width=batch_width): + _run_dummy_cuda_ipc_stream_readiness_case( + self.store, + batch_width=batch_width, + explicit_sync=explicit_sync, + sleep_cycles=sleep_cycles, + ) + + def test_dummy_cuda_ipc_put_tensor_waits_for_current_stream(self): + self._run_dummy_cuda_ipc_stream_readiness_regression(batch_width=1) + + def test_dummy_cuda_ipc_batch_put_tensor_waits_for_current_stream(self): + self._run_dummy_cuda_ipc_stream_readiness_regression(batch_width=2) + def test_00_mixed_put_and_put_tensor_concurrency(self): """Regression test for regular put and tensor put sharing dummy SHM.""" if torch is None: From 61d1403aad40c38c1cd707a3f0579639016f56f9 Mon Sep 17 00:00:00 2001 From: tancz <544463199@qq.com> Date: Wed, 12 Aug 2026 11:55:59 +0800 Subject: [PATCH 037/483] [Store] Support configurable fileread worker pool size (#3348) Make the number of FilereadWorkerPool worker threads configurable via the MC_FILEREAD_WORKERS environment variable, matching the existing MC_NOF_WORKERS pattern. Falls back to the previous default (10) when the variable is unset, empty, or invalid, so behavior is unchanged by default. Signed-off-by: tan changzhi <544463199@qq.com> --- mooncake-store/src/transfer_task.cpp | 57 ++++++++++++++++------------ 1 file changed, 33 insertions(+), 24 deletions(-) diff --git a/mooncake-store/src/transfer_task.cpp b/mooncake-store/src/transfer_task.cpp index e4b55f1822..37a147c193 100644 --- a/mooncake-store/src/transfer_task.cpp +++ b/mooncake-store/src/transfer_task.cpp @@ -19,6 +19,26 @@ #include "spdk/spdk_wrapper.h" #endif +static int GetPositiveEnvOrDefault(const char* name, int default_value) { + const char* raw_value = std::getenv(name); + if (!raw_value || raw_value[0] == '\0') { + return default_value; + } + + errno = 0; + char* end_ptr = nullptr; + long parsed = std::strtol(raw_value, &end_ptr, 10); + if (errno != 0 || end_ptr == raw_value || + (end_ptr != nullptr && *end_ptr != '\0') || parsed <= 0 || + parsed > std::numeric_limits::max()) { + LOG(WARNING) << "Invalid value for " << name << ": " << raw_value + << ", using default " << default_value; + return default_value; + } + + return static_cast(parsed); +} + #ifdef USE_NOF static bool IsTruthyEnv(const char* value) { if (!value) { @@ -54,26 +74,6 @@ static int GetSpdkNofDebugIntervalMs() { return interval_ms; } -static int GetPositiveEnvOrDefault(const char* name, int default_value) { - const char* raw_value = std::getenv(name); - if (!raw_value || raw_value[0] == '\0') { - return default_value; - } - - errno = 0; - char* end_ptr = nullptr; - long parsed = std::strtol(raw_value, &end_ptr, 10); - if (errno != 0 || end_ptr == raw_value || - (end_ptr != nullptr && *end_ptr != '\0') || parsed <= 0 || - parsed > std::numeric_limits::max()) { - LOG(WARNING) << "Invalid value for " << name << ": " << raw_value - << ", using default " << default_value; - return default_value; - } - - return static_cast(parsed); -} - static int GetSpdkNofSubmitChunkBytes() { static const int value = GetPositiveEnvOrDefault( "MC_NOF_SUBMIT_CHUNK_BYTES", mooncake::kDefaultSpdkNofSubmitChunkBytes); @@ -178,14 +178,23 @@ SpdkNofQos::SpdkNofQos(uint32_t block_size) { // threads. constexpr int kDefaultFilereadWorkers = 10; +// The number of fileread workers can be tuned via the MC_FILEREAD_WORKERS +// environment variable. Falls back to kDefaultFilereadWorkers when unset, +// empty, or invalid. +static int GetFilereadWorkerCount() { + static const int value = + GetPositiveEnvOrDefault("MC_FILEREAD_WORKERS", kDefaultFilereadWorkers); + return value; +} + FilereadWorkerPool::FilereadWorkerPool(std::shared_ptr& backend) : shutdown_(false) { - VLOG(1) << "Creating FilereadWorkerPool with " << kDefaultFilereadWorkers - << " workers"; + const int num_workers = GetFilereadWorkerCount(); + VLOG(1) << "Creating FilereadWorkerPool with " << num_workers << " workers"; // Start worker threads - workers_.reserve(kDefaultFilereadWorkers); - for (int i = 0; i < kDefaultFilereadWorkers; ++i) { + workers_.reserve(num_workers); + for (int i = 0; i < num_workers; ++i) { workers_.emplace_back(&FilereadWorkerPool::workerThread, this); } backend_ = backend; From 51e594d3a21660bdf2f6f1f11ec544b7cfb06932 Mon Sep 17 00:00:00 2001 From: SongOf <46475785+SongOf@users.noreply.github.com> Date: Wed, 12 Aug 2026 12:00:32 +0800 Subject: [PATCH 038/483] [Bugfix][TE] Install shutdown handlers before unblocking signals (#3390) Co-authored-by: maxlisongsong --- .../src/graceful_shutdown.cpp | 47 +++++++++++-------- 1 file changed, 27 insertions(+), 20 deletions(-) diff --git a/mooncake-transfer-engine/src/graceful_shutdown.cpp b/mooncake-transfer-engine/src/graceful_shutdown.cpp index 39187add49..dd77b5b44b 100644 --- a/mooncake-transfer-engine/src/graceful_shutdown.cpp +++ b/mooncake-transfer-engine/src/graceful_shutdown.cpp @@ -95,7 +95,10 @@ void signalWatcher() { _Exit(1); } -bool startSignalWatcherLocked(pid_t current_pid) { +bool startSignalWatcherLocked(pid_t current_pid, sigset_t* saved_mask, + bool* mask_blocked) { + *mask_blocked = false; + if (g_signal_pipe[0] >= 0) close(g_signal_pipe[0]); if (g_signal_pipe[1] >= 0) close(g_signal_pipe[1]); g_signal_pipe[0] = -1; @@ -107,16 +110,16 @@ bool startSignalWatcherLocked(pid_t current_pid) { // The watcher must never take SIGTERM/SIGINT itself: the handler writes // to the pipe and then pauses forever, so if it ran on the watcher // thread no reader would be left and cleanup would never run. Block both - // signals around thread creation so the watcher inherits the mask, and - // restore the caller's original mask on every path (it may have had its - // own reasons to block these signals). + // signals before creating the watcher so it inherits the mask; the caller + // restores the original mask once sigaction() has installed the handlers, + // so a signal arriving mid-install is not delivered while the disposition + // is still SIG_DFL. sigset_t watcher_block_set; - sigset_t saved_mask; sigemptyset(&watcher_block_set); sigaddset(&watcher_block_set, SIGTERM); sigaddset(&watcher_block_set, SIGINT); - const bool mask_blocked = - pthread_sigmask(SIG_BLOCK, &watcher_block_set, &saved_mask) == 0; + *mask_blocked = + pthread_sigmask(SIG_BLOCK, &watcher_block_set, saved_mask) == 0; bool watcher_started = false; try { @@ -128,8 +131,6 @@ bool startSignalWatcherLocked(pid_t current_pid) { } catch (...) { } - if (mask_blocked) pthread_sigmask(SIG_SETMASK, &saved_mask, nullptr); - if (!watcher_started) { close(g_signal_pipe[0]); close(g_signal_pipe[1]); @@ -189,18 +190,24 @@ void installGracefulShutdownHandlers() { g_atexit_registered = true; } - if (!startSignalWatcherLocked(current_pid)) return; - - struct sigaction sa{}; - sa.sa_handler = shutdownSignalHandler; - sigemptyset(&sa.sa_mask); - sigaddset(&sa.sa_mask, SIGTERM); - sigaddset(&sa.sa_mask, SIGINT); - sa.sa_flags = 0; - sigaction(SIGTERM, &sa, nullptr); - sigaction(SIGINT, &sa, nullptr); + sigset_t saved_mask; + bool mask_blocked = false; + if (startSignalWatcherLocked(current_pid, &saved_mask, &mask_blocked)) { + struct sigaction sa{}; + sa.sa_handler = shutdownSignalHandler; + sigemptyset(&sa.sa_mask); + sigaddset(&sa.sa_mask, SIGTERM); + sigaddset(&sa.sa_mask, SIGINT); + sa.sa_flags = 0; + sigaction(SIGTERM, &sa, nullptr); + sigaction(SIGINT, &sa, nullptr); + + g_handlers_installed = true; + } - g_handlers_installed = true; + // Restore last: a signal that arrived mid-install stays pending until the + // handlers above are in place. + if (mask_blocked) pthread_sigmask(SIG_SETMASK, &saved_mask, nullptr); } } // namespace mooncake From d6ccd4ab485263cf1f031c22738a1b4e00f8b110 Mon Sep 17 00:00:00 2001 From: Icedcoco <102317026+Icedcoco@users.noreply.github.com> Date: Wed, 12 Aug 2026 12:32:28 +0800 Subject: [PATCH 039/483] [Store] Verify supervisor view propagation (#3352) Co-authored-by: Yuchen Kou --- .../ha/leadership/high_availability_test.cpp | 87 +++++++++++++++++++ 1 file changed, 87 insertions(+) diff --git a/mooncake-store/tests/ha/leadership/high_availability_test.cpp b/mooncake-store/tests/ha/leadership/high_availability_test.cpp index dbd2c5fe97..f54a87d212 100644 --- a/mooncake-store/tests/ha/leadership/high_availability_test.cpp +++ b/mooncake-store/tests/ha/leadership/high_availability_test.cpp @@ -14,6 +14,7 @@ #endif #include "ha/leadership/leader_coordinator_factory.h" #include "ha/leadership/high_availability_test_fixture.h" +#include "master_service.h" #include "types.h" namespace mooncake { @@ -95,8 +96,94 @@ std::unique_ptr CreateEtcdCoordinatorOrNull( return std::move(coordinator.value()); } +class FakeLeaderCoordinator : public ha::LeaderCoordinator { + public: + explicit FakeLeaderCoordinator(ViewVersionId view_version) + : session_{.view = {.leader_address = "fake-leader", + .view_version = view_version}, + .owner_token = "fake-owner", + .lease_ttl = std::chrono::milliseconds(0)} {} + + tl::expected, ErrorCode> ReadCurrentView() + override { + return std::optional{}; + } + + tl::expected TryAcquireLeadership( + const std::string& /*leader_address*/) override { + return ha::AcquireLeadershipResult{ + .status = ha::AcquireLeadershipStatus::ACQUIRED, + .session = session_, + .observed_view = std::nullopt, + }; + } + + tl::expected RenewLeadership( + const ha::LeadershipSession& /*session*/) override { + return true; + } + + tl::expected WaitForViewChange( + std::optional /*known_version*/, + std::chrono::milliseconds /*timeout*/) override { + return ha::ViewChangeResult{}; + } + + tl::expected, ErrorCode> + StartLeadershipMonitor( + const ha::LeadershipSession& /*session*/, + ha::LeadershipLostCallback /*on_leadership_lost*/) override { + return tl::make_unexpected(ErrorCode::UNAVAILABLE_IN_CURRENT_STATUS); + } + + ErrorCode ReleaseLeadership( + const ha::LeadershipSession& /*session*/) override { + return ErrorCode::OK; + } + + private: + ha::LeadershipSession session_; +}; + } // namespace +TEST_F(HighAvailabilityTest, AcquiredViewFlowsIntoServingMasterService) { + constexpr ViewVersionId kAcquiredView = 42; + FakeLeaderCoordinator coordinator(kAcquiredView); + auto acquired = coordinator.TryAcquireLeadership("primary"); + ASSERT_TRUE(acquired.has_value()); + ASSERT_TRUE(acquired->session.has_value()); + + MasterServiceSupervisorConfig supervisor_config; + supervisor_config.default_kv_lease_ttl = DEFAULT_DEFAULT_KV_LEASE_TTL; + supervisor_config.default_kv_soft_pin_ttl = DEFAULT_KV_SOFT_PIN_TTL_MS; + supervisor_config.allow_evict_soft_pinned_objects = true; + supervisor_config.enable_metric_reporting = false; + supervisor_config.metrics_port = 0; + supervisor_config.eviction_ratio = DEFAULT_EVICTION_RATIO; + supervisor_config.eviction_high_watermark_ratio = + DEFAULT_EVICTION_HIGH_WATERMARK_RATIO; + supervisor_config.nof_eviction_ratio = DEFAULT_NOF_EVICTION_RATIO; + supervisor_config.nof_eviction_high_watermark_ratio = + DEFAULT_NOF_EVICTION_HIGH_WATERMARK_RATIO; + supervisor_config.client_live_ttl_sec = DEFAULT_CLIENT_LIVE_TTL_SEC; + supervisor_config.nof_heartbeat_interval_sec = + DEFAULT_NOF_HEARTBEAT_INTERVAL_SEC; + supervisor_config.nof_heartbeat_probe_timeout_ms = + DEFAULT_NOF_HEARTBEAT_PROBE_TIMEOUT_MS; + supervisor_config.nof_heartbeat_failures_threshold = + DEFAULT_NOF_HEARTBEAT_FAILURES_THRESHOLD; + supervisor_config.enable_offload = false; + + WrappedMasterServiceConfig serving_config( + supervisor_config, acquired->session->view.view_version); + MasterService service{MasterServiceConfig(serving_config)}; + + auto ping = service.Ping(generate_uuid()); + ASSERT_TRUE(ping.has_value()); + EXPECT_EQ(kAcquiredView, ping->view_version_id); +} + #ifdef STORE_USE_ETCD TEST_F(HighAvailabilityTest, EtcdBasicOperations) { From f87babd4367a94de8d1e9d545305c8386dc63c1a Mon Sep 17 00:00:00 2001 From: Cruz Zhao Date: Wed, 12 Aug 2026 12:41:15 +0800 Subject: [PATCH 040/483] [TENT] drain control callbacks on unregister (#3370) * fix(tent): drain control callbacks on unregister * style(tent): format callback drain waits * fix(tent): unregister rdma bootstrap callback on uninstall * style(tent): format callback invocation guard --- .../tent/include/tent/runtime/control_plane.h | 26 ++- .../tent/src/runtime/control_plane.cpp | 149 +++++++++++++++++- .../src/transport/rdma/rdma_transport.cpp | 6 + 3 files changed, 169 insertions(+), 12 deletions(-) diff --git a/mooncake-transfer-engine/tent/include/tent/runtime/control_plane.h b/mooncake-transfer-engine/tent/include/tent/runtime/control_plane.h index a4739f27aa..f8d2c92f9e 100644 --- a/mooncake-transfer-engine/tent/include/tent/runtime/control_plane.h +++ b/mooncake-transfer-engine/tent/include/tent/runtime/control_plane.h @@ -19,9 +19,12 @@ #include #include +#include +#include #include #include #include +#include #include #include #include @@ -121,13 +124,9 @@ class ControlService { SegmentManager& segmentManager() { return *manager_.get(); } - void setBootstrapRdmaCallback(const OnReceiveBootstrap& callback) { - bootstrap_callback_ = callback; - } + void setBootstrapRdmaCallback(const OnReceiveBootstrap& callback); - void setNotifyCallback(const OnNotify& callback) { - notify_callback_ = callback; - } + void setNotifyCallback(const OnNotify& callback); Status start(uint16_t& port, bool ipv6_ = false); @@ -160,12 +159,27 @@ class ControlService { void onSegmentUpdated(const std::string_view& request, std::string& response); + void finishBootstrapCallback(); + + void finishNotifyCallback(); + private: std::unique_ptr manager_; std::shared_ptr rpc_server_; + std::mutex bootstrap_cb_mutex_; + std::condition_variable bootstrap_cb_cv_; + size_t bootstrap_callbacks_in_flight_ = 0; + std::chrono::milliseconds callback_drain_timeout_{std::chrono::seconds(5)}; OnReceiveBootstrap bootstrap_callback_; + static thread_local const ControlService* active_bootstrap_service_; + + std::mutex notify_cb_mutex_; + std::condition_variable notify_cb_cv_; + size_t notify_callbacks_in_flight_ = 0; OnNotify notify_callback_; + static thread_local const ControlService* active_notify_service_; + TransferEngineImpl* impl_; }; diff --git a/mooncake-transfer-engine/tent/src/runtime/control_plane.cpp b/mooncake-transfer-engine/tent/src/runtime/control_plane.cpp index ca5facb0e6..c57e07e5f5 100644 --- a/mooncake-transfer-engine/tent/src/runtime/control_plane.cpp +++ b/mooncake-transfer-engine/tent/src/runtime/control_plane.cpp @@ -17,6 +17,7 @@ #include #include +#include #include "tent/common/status.h" #include "tent/common/utils/os.h" @@ -25,6 +26,29 @@ namespace mooncake { namespace tent { +namespace { + +template +class CallbackInvocationGuard { + public: + explicit CallbackInvocationGuard(Fn on_exit) + : on_exit_(std::move(on_exit)) {} + + ~CallbackInvocationGuard() { on_exit_(); } + + CallbackInvocationGuard(const CallbackInvocationGuard&) = delete; + CallbackInvocationGuard& operator=(const CallbackInvocationGuard&) = delete; + + private: + Fn on_exit_; +}; + +} // namespace + +thread_local const ControlService* ControlService::active_bootstrap_service_ = + nullptr; +thread_local const ControlService* ControlService::active_notify_service_ = + nullptr; thread_local CoroRpcAgent tl_rpc_agent; Status ControlClient::getSegmentDesc(const std::string& server_addr, @@ -218,7 +242,62 @@ ControlService::ControlService(const std::string& type, }); } -ControlService::~ControlService() {} +ControlService::~ControlService() { + // Stop RPC workers while callback state and synchronization primitives are + // still alive. Member destruction would otherwise tear them down first. + rpc_server_.reset(); +} + +void ControlService::setBootstrapRdmaCallback( + const OnReceiveBootstrap& callback) { + std::unique_lock guard(bootstrap_cb_mutex_); + if (active_bootstrap_service_ == this) { + bootstrap_callback_ = callback; + return; + } + bootstrap_callback_ = nullptr; + if (!bootstrap_cb_cv_.wait_for(guard, callback_drain_timeout_, [this] { + return bootstrap_callbacks_in_flight_ == 0; + })) { + LOG(ERROR) + << "Timed out waiting for BootstrapRdma callbacks to drain, " + << "in_flight=" << bootstrap_callbacks_in_flight_ + << ", timeout_ms=" << callback_drain_timeout_.count() + << ". Continue replacing the callback to keep shutdown bounded."; + } + bootstrap_callback_ = callback; +} + +void ControlService::setNotifyCallback(const OnNotify& callback) { + std::unique_lock guard(notify_cb_mutex_); + if (active_notify_service_ == this) { + notify_callback_ = callback; + return; + } + notify_callback_ = nullptr; + if (!notify_cb_cv_.wait_for(guard, callback_drain_timeout_, [this] { + return notify_callbacks_in_flight_ == 0; + })) { + LOG(ERROR) << "Timed out waiting for Notify callbacks to drain, " + << "in_flight=" << notify_callbacks_in_flight_ + << ", timeout_ms=" << callback_drain_timeout_.count() + << ". Continue replacing the callback to keep shutdown " + << "bounded."; + } + notify_callback_ = callback; +} + +void ControlService::finishBootstrapCallback() { + std::lock_guard guard(bootstrap_cb_mutex_); + --bootstrap_callbacks_in_flight_; + bootstrap_cb_cv_.notify_all(); +} + +void ControlService::finishNotifyCallback() { + std::lock_guard guard(notify_cb_mutex_); + --notify_callbacks_in_flight_; + notify_cb_cv_.notify_all(); +} Status ControlService::start(uint16_t& port, bool ipv6_) { return rpc_server_->start(port, ipv6_); @@ -234,10 +313,46 @@ void ControlService::onGetSegmentDesc(const std::string_view& request, void ControlService::onBootstrapRdma(const std::string_view& request, std::string& response) { std::string mutable_request(request); - BootstrapDesc request_desc = - json::parse(std::string(request)).get(); + OnReceiveBootstrap callback; + { + std::lock_guard guard(bootstrap_cb_mutex_); + if (!bootstrap_callback_) { + BootstrapDesc response_desc; + response_desc.reply_msg = "NOT_READY: transport not initialized"; + json j = response_desc; + response = j.dump(); + return; + } + callback = bootstrap_callback_; + ++bootstrap_callbacks_in_flight_; + } + BootstrapDesc response_desc; - if (bootstrap_callback_) bootstrap_callback_(request_desc, response_desc); + { + const ControlService* previous_service = active_bootstrap_service_; + active_bootstrap_service_ = this; + CallbackInvocationGuard invocation_guard([this, previous_service] { + active_bootstrap_service_ = previous_service; + finishBootstrapCallback(); + }); + + try { + BootstrapDesc request_desc = + json::parse(mutable_request).get(); + int rc = callback(request_desc, response_desc); + if (rc != 0 && response_desc.reply_msg.empty()) { + response_desc.reply_msg = "BootstrapRdma callback failed"; + } + } catch (const std::exception& e) { + LOG(ERROR) << "onBootstrapRdma failed: " << e.what(); + response_desc.reply_msg = + std::string("BootstrapRdma callback failed: ") + e.what(); + } catch (...) { + LOG(ERROR) << "onBootstrapRdma failed with unknown exception"; + response_desc.reply_msg = "BootstrapRdma callback failed"; + } + } + json j = response_desc; response = j.dump(); } @@ -295,8 +410,30 @@ void ControlService::onRecvData(const std::string_view& request, void ControlService::onNotify(const std::string_view& request, std::string& response) { - Notification message = json::parse(request).get(); - if (notify_callback_) notify_callback_(message); + (void)response; + OnNotify callback; + { + std::lock_guard guard(notify_cb_mutex_); + if (!notify_callback_) return; + callback = notify_callback_; + ++notify_callbacks_in_flight_; + } + + const ControlService* previous_service = active_notify_service_; + active_notify_service_ = this; + CallbackInvocationGuard invocation_guard([this, previous_service] { + active_notify_service_ = previous_service; + finishNotifyCallback(); + }); + + try { + Notification message = json::parse(request).get(); + callback(message); + } catch (const std::exception& e) { + LOG(ERROR) << "onNotify failed: " << e.what(); + } catch (...) { + LOG(ERROR) << "onNotify failed with unknown exception"; + } } void ControlService::onProbe(const std::string_view& request, diff --git a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp index f4d2e6177b..7bdb5bd1c3 100644 --- a/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/rdma/rdma_transport.cpp @@ -378,6 +378,12 @@ Status RdmaTransport::install(std::string& local_segment_name, } Status RdmaTransport::uninstall() { + // ControlService may still receive BootstrapRdma RPCs while uninstall is + // running. Unregister and drain the callback before destroying workers, + // contexts, and other state used by onSetupRdmaConnections(). Keep this + // outside installed_ so partially-installed transports are covered too. + if (metadata_) metadata_->setBootstrapRdmaCallback(nullptr); + if (installed_) { // Stop notification worker thread notify_worker_running_ = false; From d5f6ffd7e2f85113e67542ad99759808be89fd1a Mon Sep 17 00:00:00 2001 From: Feng Ren Date: Wed, 12 Aug 2026 12:50:20 +0800 Subject: [PATCH 041/483] [TransferEngine] Add instant bandwidth reporting to tebench (#3358) * benchmark: report instant transfer bandwidth * remove namespace * Update the counting of avg lat and docs --- docs/source/design/tent/tebench.md | 14 +++++- mooncake-transfer-engine/benchmark/main.cpp | 50 ++++++++++++++++++++ mooncake-transfer-engine/benchmark/utils.cpp | 10 +++- mooncake-transfer-engine/benchmark/utils.h | 2 + 4 files changed, 73 insertions(+), 3 deletions(-) diff --git a/docs/source/design/tent/tebench.md b/docs/source/design/tent/tebench.md index 4e9ce606f7..bdc87d00d4 100644 --- a/docs/source/design/tent/tebench.md +++ b/docs/source/design/tent/tebench.md @@ -75,6 +75,15 @@ On the initiator machine: --duration=5 ``` +### 3.3 Request Pacing + +Use `--request_interval_us=` to add a per-thread delay before each +transfer batch. The value is in microseconds; `0` disables pacing. + +When pacing is enabled, `Avg Lat (us)` includes the pacing gap because it is +computed from wall-clock runtime, while `Avg Tx (us)` and Tx percentiles only +measure transfer execution time. + ## 4. Output Metrics Each output row corresponds to one benchmark configuration. @@ -84,8 +93,9 @@ Each output row corresponds to one benchmark configuration. | `BlkSize (B)` | Block size per request (bytes) | | `Batch` | Number of requests per submission | | `BW (GB/S)` | Throughput (total bytes / total time) | -| `Avg Lat (us)` | Average end-to-end latency (scaled by thread count) | -| `Avg Tx (us)` | Average per-transfer execution time | +| `Avg Inst GB/s` | Average per-transfer instantaneous bandwidth | +| `Avg Lat (us)` | Average wall-clock per-operation latency, including pacing or scheduling gaps (scaled by thread count) | +| `Avg Tx (us)` | Average per-transfer execution time, excluding gaps | | `P99 Tx (us)` | P99 transfer latency | | `P999 Tx (us)` | P999 transfer latency | diff --git a/mooncake-transfer-engine/benchmark/main.cpp b/mooncake-transfer-engine/benchmark/main.cpp index ad756e980b..f3935d8397 100644 --- a/mooncake-transfer-engine/benchmark/main.cpp +++ b/mooncake-transfer-engine/benchmark/main.cpp @@ -22,8 +22,22 @@ #include "tent_backend.h" #endif +#include +#include +#include + using namespace mooncake::tent; +uint64_t steadyClockNs() { + const auto now = std::chrono::steady_clock::now().time_since_epoch(); + return std::chrono::duration_cast(now).count(); +} + +double gbPerSecond(uint64_t bytes, double duration_us) { + if (duration_us <= 0.0) return 0.0; + return static_cast(bytes) / (1000.0 * duration_us); +} + int processBatchSizes( BenchRunner& runner, size_t block_size, size_t batch_size, int num_threads, const std::vector& qos_classes, @@ -46,6 +60,17 @@ int processBatchSizes( XferBenchStats tight_stats; XferBenchStats loose_stats; std::mutex mutex; + std::atomic measurement_ready{0}; + std::atomic measurement_started{false}; + auto paceRequest = [&]() { + if (XferBenchConfig::request_interval_us == 0) return; + const uint64_t interval_ns = + XferBenchConfig::request_interval_us * 1000ull; + const uint64_t target_ns = steadyClockNs() + interval_ns; + while (steadyClockNs() < target_ns) { + std::this_thread::yield(); + } + }; size_t address_stride_bytes = XferBenchConfig::max_block_size * XferBenchConfig::max_batch_size; if (!workload_classes.empty()) { @@ -94,24 +119,41 @@ int processBatchSizes( thread_batch_size, opcode, deadlineNs(), intent_type); } + if (measurement_ready.fetch_add(1, std::memory_order_acq_rel) + 1 == + num_threads) { + measurement_started.store(true, std::memory_order_release); + } else { + while (!measurement_started.load(std::memory_order_acquire)) { + std::this_thread::yield(); + } + } timer.reset(); std::vector transfer_duration; + std::vector thread_instant_bandwidth; if (mixed_opcode) { while (timer.lap_us(false) < XferBenchConfig::duration * 1000000ull) { + const uint64_t batch_bytes = + thread_block_size * thread_batch_size; uint8_t pattern = 0; if (XferBenchConfig::check_consistency) pattern = fillData((void*)local_addr, thread_block_size * thread_batch_size); + paceRequest(); auto val = runner.runSingleTransfer( local_addr, target_addr, thread_block_size, thread_batch_size, WRITE, deadlineNs(), intent_type); + thread_instant_bandwidth.push_back( + gbPerSecond(batch_bytes, val)); transfer_duration.push_back(val); fillData((void*)local_addr, thread_block_size * thread_batch_size); + paceRequest(); val = runner.runSingleTransfer( local_addr, target_addr, thread_block_size, thread_batch_size, READ, deadlineNs(), intent_type); + thread_instant_bandwidth.push_back( + gbPerSecond(batch_bytes, val)); if (XferBenchConfig::check_consistency) verifyData((void*)local_addr, thread_block_size * thread_batch_size, pattern); @@ -120,9 +162,14 @@ int processBatchSizes( } else { while (timer.lap_us(false) < XferBenchConfig::duration * 1000000ull) { + const uint64_t batch_bytes = + thread_block_size * thread_batch_size; + paceRequest(); auto val = runner.runSingleTransfer( local_addr, target_addr, thread_block_size, thread_batch_size, opcode, deadlineNs(), intent_type); + thread_instant_bandwidth.push_back( + gbPerSecond(batch_bytes, val)); transfer_duration.push_back(val); } } @@ -130,9 +177,12 @@ int processBatchSizes( std::lock_guard lock(mutex); stats.total_duration.add(total_duration); stats.transfer_duration.add(transfer_duration); + stats.instant_bandwidth.add(thread_instant_bandwidth); if (qos_enabled) { qos_stats[qos_class].total_duration.add(total_duration); qos_stats[qos_class].transfer_duration.add(transfer_duration); + qos_stats[qos_class].instant_bandwidth.add( + thread_instant_bandwidth); } auto& group_stats = tight ? tight_stats : loose_stats; group_stats.total_duration.add(total_duration); diff --git a/mooncake-transfer-engine/benchmark/utils.cpp b/mooncake-transfer-engine/benchmark/utils.cpp index cbc839d3ac..88cb49d0b3 100644 --- a/mooncake-transfer-engine/benchmark/utils.cpp +++ b/mooncake-transfer-engine/benchmark/utils.cpp @@ -60,6 +60,9 @@ DEFINE_double(qos_link_capacity_gbps, 0.0, "Link capacity in GB/s for total utilization (0 reports N/A)."); DEFINE_string(qos_output_jsonl, "", "Append versioned QoS metric records to this JSONL file."); +DEFINE_uint64(request_interval_us, 0, + "Per-thread delay before issuing each transfer batch, in " + "microseconds. 0 disables pacing."); DEFINE_uint64(deadline_us, 0, "tent only: relative per-transfer deadline in microseconds for " "tight worker threads (0 disables deadline tagging); cannot be " @@ -115,6 +118,7 @@ std::string XferBenchConfig::qos_classes_json; std::string XferBenchConfig::workload_classes_json; double XferBenchConfig::qos_link_capacity_gbps = 0.0; std::string XferBenchConfig::qos_output_jsonl; +uint64_t XferBenchConfig::request_interval_us = 0; uint64_t XferBenchConfig::deadline_us = 0; int XferBenchConfig::deadline_tight_threads = 0; bool XferBenchConfig::deadline_bw_arbitration = false; @@ -151,6 +155,7 @@ void XferBenchConfig::loadFromFlags() { workload_classes_json = FLAGS_workload_classes_json; qos_link_capacity_gbps = FLAGS_qos_link_capacity_gbps; qos_output_jsonl = FLAGS_qos_output_jsonl; + request_interval_us = FLAGS_request_interval_us; deadline_us = FLAGS_deadline_us; deadline_tight_threads = FLAGS_deadline_tight_threads; deadline_bw_arbitration = FLAGS_deadline_bw_arbitration; @@ -191,7 +196,8 @@ void printStatsHeader() { std::cout << std::left << std::setw(14) << "BlkSize (B)" << std::setw(8) << "Batch" - << std::setw(14) << "BW (GB/S)" + << std::setw(14) << "BW (GB/s)" + << std::setw(18) << "Avg Inst GB/s" << std::setw(14) << "Avg Lat (us)" << std::setw(14) << "Avg Tx (us)" << std::setw(14) << "P99 Tx (us)" @@ -211,6 +217,7 @@ void printStats(size_t block_size, size_t batch_size, XferBenchStats& stats, avg_latency = (total_duration * num_threads / num_ops); throughput_gb = (((double)total_data_transferred / (1000 * 1000 * 1000)) / (total_duration / 1e6)); // In GB/Sec + const double avg_instant_gbps = stats.instant_bandwidth.avg(); // Tabulate print with fixed width for each string // clang-format off @@ -218,6 +225,7 @@ void printStats(size_t block_size, size_t batch_size, XferBenchStats& stats, << std::setw(14) << block_size << std::setw(8) << batch_size << std::setw(14) << throughput_gb + << std::setw(18) << avg_instant_gbps << std::setprecision(1) << std::setw(14) << avg_latency << std::setw(14) << stats.transfer_duration.avg() diff --git a/mooncake-transfer-engine/benchmark/utils.h b/mooncake-transfer-engine/benchmark/utils.h index 9d6cab1e6c..5dc9833dd6 100644 --- a/mooncake-transfer-engine/benchmark/utils.h +++ b/mooncake-transfer-engine/benchmark/utils.h @@ -76,6 +76,7 @@ struct XferBenchConfig { static std::string workload_classes_json; static double qos_link_capacity_gbps; static std::string qos_output_jsonl; + static uint64_t request_interval_us; static uint64_t deadline_us; static int deadline_tight_threads; static bool deadline_bw_arbitration; @@ -147,6 +148,7 @@ struct XferMetricStats { struct XferBenchStats { XferMetricStats total_duration; XferMetricStats transfer_duration; + XferMetricStats instant_bandwidth; }; class XferBenchTimer { From e6ed919fe73b83a673d8ed54695cf8966010afc3 Mon Sep 17 00:00:00 2001 From: mjwtom Date: Wed, 12 Aug 2026 13:21:14 +0800 Subject: [PATCH 042/483] [Bugfix][Store] Update bucket timestamp after offload to avoid immediate eviction (#3097) (#3101) BatchOffload now updates the bucket's last_access_ns_ with the current time when committing metadata under LRU eviction policy. This prevents freshly offloaded buckets from being evicted immediately due to having a zero timestamp. Co-authored-by: majingwei --- mooncake-store/include/storage_backend.h | 20 ++++ mooncake-store/src/storage_backend.cpp | 17 ++- mooncake-store/tests/storage_backend_test.cpp | 110 ++++++++++++++++++ 3 files changed, 145 insertions(+), 2 deletions(-) diff --git a/mooncake-store/include/storage_backend.h b/mooncake-store/include/storage_backend.h index 14e8a822af..b6acb75a6a 100644 --- a/mooncake-store/include/storage_backend.h +++ b/mooncake-store/include/storage_backend.h @@ -186,6 +186,20 @@ enum class BucketEvictionPolicy { LRU, // Evict least recently read bucket first }; +inline std::ostream& operator<<(std::ostream& os, + const BucketEvictionPolicy& policy) { + switch (policy) { + case BucketEvictionPolicy::NONE: + return os << "none"; + case BucketEvictionPolicy::FIFO: + return os << "fifo"; + case BucketEvictionPolicy::LRU: + return os << "lru"; + default: + return os << "unknown"; + } +} + struct BucketBackendConfig { int64_t bucket_size_limit = 256 * kMB; // Max total size of a single bucket (256 MB) @@ -1107,6 +1121,12 @@ class BucketStorageBackend : public StorageBackendInterface { tl::expected, ErrorCode> GetFileInstance() const; + // Test-only: number of entries in the LRU eviction index. + size_t GetLruIndexSizeForTest() const { + SharedMutexLocker lock(&mutex_, shared_lock); + return lru_index_.size(); + } + private: // Alignment helper functions for O_DIRECT I/O static constexpr size_t kDirectIOAlignment = 4096; diff --git a/mooncake-store/src/storage_backend.cpp b/mooncake-store/src/storage_backend.cpp index d206571402..5a5e1588db 100644 --- a/mooncake-store/src/storage_backend.cpp +++ b/mooncake-store/src/storage_backend.cpp @@ -1770,6 +1770,7 @@ tl::expected BucketStorageBackend::BatchOffload( ReleasePreparedWrite(pending); return tl::make_unexpected(write_bucket_result.error()); } + VLOG(1) << "Written bucket with id: " << bucket_id; // Save a copy of bucket->keys before std::move(bucket) into buckets_ // consumes the shared_ptr. Needed for complete_handler and rollback. const auto bucket_keys = bucket->keys; @@ -1803,8 +1804,16 @@ tl::expected BucketStorageBackend::BatchOffload( CHECK(inserted) << "Reserved key became duplicated: " << bucket_keys[i]; } + auto ts = 0LL; + // Update LRU timestamp for in case of eviction. + if (bucket_backend_config_.eviction_policy == + BucketEvictionPolicy::LRU) { + ts = + std::chrono::steady_clock::now().time_since_epoch().count(); + bucket->last_access_ns_.store(ts, std::memory_order_relaxed); + } buckets_.emplace(bucket_id, std::move(bucket)); - lru_index_.emplace(0LL, bucket_id); + lru_index_.emplace(ts, bucket_id); } if (duplicate_found) { LOG(ERROR) << "Reserved key became duplicated before commit, " @@ -2731,7 +2740,9 @@ void BucketStorageBackend::RollbackCommittedBucket( // Remove bucket metadata total_size_ -= bucket_meta->meta_size; - lru_index_.erase({0LL, bucket_id}); + lru_index_.erase( + {bucket_meta->last_access_ns_.load(std::memory_order_relaxed), + bucket_id}); buckets_.erase(bucket_it); } @@ -3127,6 +3138,8 @@ tl::expected BucketStorageBackend::FinalizeEviction( if (bucket_cleanup_failed) { cleanup_failed_count++; } + VLOG(1) << "Evicted bucket with id: " << bucket_id + << ", policy: " << bucket_backend_config_.eviction_policy; } if (!pending.buckets.empty()) { LOG(INFO) << "[Evict] finalized: attempted=" << pending.buckets.size() diff --git a/mooncake-store/tests/storage_backend_test.cpp b/mooncake-store/tests/storage_backend_test.cpp index 1f3e1ae246..605b5c4877 100644 --- a/mooncake-store/tests/storage_backend_test.cpp +++ b/mooncake-store/tests/storage_backend_test.cpp @@ -4725,4 +4725,114 @@ TEST_F(StorageBackendTest, std::make_optional(value_b)); } +TEST_F(StorageBackendTest, + BucketStorageBackend_OffloadSetsBucketAccessTimestamp) { + // Regression test for "update bucket ts after offload": BatchOffload must + // stamp the new bucket's last_access_ns_ with the current time when + // committing metadata. The timestamp itself is private, so we verify it + // through LRU eviction behavior: + // - Bucket A is offloaded and then read, so its timestamp is t1 > 0. + // - Bucket B is offloaded afterwards; with the fix its timestamp is + // t2 > t1, without the fix it stays 0. + // - Offloading bucket C overflows the quota and evicts exactly one + // bucket. The LRU victim must be A (oldest timestamp). If B's + // timestamp were 0, B would be evicted immediately instead. + // Additionally, a failing complete_handler must roll back the committed + // bucket, including its lru_index_ entry (checked via + // GetLruIndexSizeForTest at every stage). + FileStorageConfig config; + config.storage_filepath = data_path; + BucketBackendConfig bucket_config; + bucket_config.eviction_policy = BucketEvictionPolicy::LRU; + // Each bucket holds one 64 KB object (plus a small metadata blob); the + // quota fits two buckets but not three, so the third offload evicts + // exactly one bucket. + constexpr size_t kValueSize = 64 * 1024; + bucket_config.max_total_size = 150 * 1024; + BucketStorageBackend storage_backend(config, bucket_config); + ASSERT_TRUE(storage_backend.Init()); + + auto offload_one = [&](const std::string& key, + std::vector* evicted_keys) { + auto buf = std::make_unique(kValueSize); + std::memset(buf.get(), 'x', kValueSize); + std::unordered_map> batch; + batch.emplace(key, std::vector{Slice{buf.get(), kValueSize}}); + auto result = storage_backend.BatchOffload( + batch, + [](const std::vector&, + std::vector&) { return ErrorCode::OK; }, + [evicted_keys](const std::vector& keys) + -> tl::expected { + if (evicted_keys != nullptr) { + evicted_keys->insert(evicted_keys->end(), keys.begin(), + keys.end()); + } + return {}; + }); + ASSERT_TRUE(result.has_value()) << "Offload failed for key " << key; + }; + + // Bucket A: offload, then read so its access timestamp becomes t1 > 0. + offload_one("lru_ts_key_a", nullptr); + EXPECT_EQ(storage_backend.GetLruIndexSizeForTest(), 1u); + + // Rollback on complete_handler failure: the committed bucket must be + // rolled back and its lru_index_ entry removed. Use a small value so this + // offload stays under the quota and does not trigger eviction. + { + std::string rollback_key = "lru_ts_rollback_key"; + std::string rollback_data = "rollback_data"; + std::unordered_map> batch; + batch.emplace(rollback_key, + std::vector{ + Slice{rollback_data.data(), rollback_data.size()}}); + auto rollback_res = storage_backend.BatchOffload( + batch, [](const std::vector&, + std::vector&) { + return ErrorCode::INTERNAL_ERROR; + }); + ASSERT_FALSE(rollback_res.has_value()); + EXPECT_EQ(rollback_res.error(), ErrorCode::INTERNAL_ERROR); + + auto exist_res = storage_backend.IsExist(rollback_key); + ASSERT_TRUE(exist_res.has_value()); + EXPECT_FALSE(exist_res.value()) << "Rolled-back key must not exist"; + + EXPECT_EQ(storage_backend.GetLruIndexSizeForTest(), 1u) + << "Rolled-back bucket's lru_index_ entry must be removed"; + } + + { + auto read_buf = std::make_unique(kValueSize); + std::unordered_map load_slices; + load_slices.emplace("lru_ts_key_a", Slice{read_buf.get(), kValueSize}); + ASSERT_TRUE(storage_backend.BatchLoad(load_slices).has_value()); + } + std::this_thread::sleep_for(std::chrono::milliseconds(2)); + + // Bucket B: freshly offloaded. Its timestamp must be t2 > t1 (it would + // be 0 without the timestamp update in BatchOffload). + offload_one("lru_ts_key_b", nullptr); + EXPECT_EQ(storage_backend.GetLruIndexSizeForTest(), 2u); + + // Bucket C: overflows the quota and forces eviction of one bucket. + std::vector evicted_keys; + offload_one("lru_ts_key_c", &evicted_keys); + EXPECT_EQ(storage_backend.GetLruIndexSizeForTest(), 2u) + << "lru_index_ must stay in sync with buckets_ after eviction"; + + ASSERT_EQ(evicted_keys.size(), 1u) + << "Exactly one bucket should be evicted"; + EXPECT_EQ(evicted_keys[0], "lru_ts_key_a") + << "LRU victim should be the previously read bucket A, not the " + "freshly offloaded bucket B"; + + auto exist_b = storage_backend.IsExist("lru_ts_key_b"); + ASSERT_TRUE(exist_b.has_value()); + EXPECT_TRUE(exist_b.value()) + << "Freshly offloaded bucket must not be evicted immediately " + "(its access timestamp must not be 0)"; +} + } // namespace mooncake::test From 67aaf6de74f35c7c5cfd072e8643b0a80eb5999a Mon Sep 17 00:00:00 2001 From: Aoi Date: Wed, 12 Aug 2026 15:47:36 +0800 Subject: [PATCH 043/483] [Store] Add deterministic MasterService eviction scenarios (#3327) --- mooncake-store/include/master_service.h | 9 +- mooncake-store/tests/CMakeLists.txt | 3 +- mooncake-store/tests/batch_evict_test.cpp | 468 -------------- .../tests/ha/master_service_ha_test.cpp | 88 --- mooncake-store/tests/master_scenario.cpp | 247 +++++++- mooncake-store/tests/master_scenario.h | 290 +++++++++ mooncake-store/tests/master_scenario_test.cpp | 50 +- .../master_service_evict_scenario_test.cpp | 592 ++++++++++++++++++ 8 files changed, 1173 insertions(+), 574 deletions(-) delete mode 100644 mooncake-store/tests/batch_evict_test.cpp create mode 100644 mooncake-store/tests/master_service_evict_scenario_test.cpp diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index cd2d357100..4e4270d75d 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -81,10 +81,7 @@ class SnapshotChildProcessTest; // exposing test-only accessors on MasterService itself. class PromotionOnHitTest; class MasterServiceTenantQuotaTest; -// Friended so the BatchEvict correctness tests can invoke the private -// BatchEvict entry point and seed lease timestamps directly, instead of -// relying on segment pressure plus the background eviction thread. -class BatchEvictTest; +class MasterScenario; class MasterServiceHATest; // Friended so the processing_keys double-erase reproduction test can // invalidate a segment allocator via PrepareUnmountSegment WITHOUT the @@ -121,7 +118,9 @@ class MasterService { friend class test::PromotionOnHitTest; friend class benchmarks::BatchEvictBench; friend class test::MasterServiceTenantQuotaTest; - friend class test::BatchEvictTest; + // The scenario DSL controls lease timestamps so eviction tests do not + // depend on sleeps or the background eviction thread. + friend class test::MasterScenario; // double-erase processing_keys UAF repro (2026-08-03 prod segfault) friend class test::MasterServiceProcessingKeyDoubleEraseTest; friend class MasterSnapshotManager; // Allow access to internal state for diff --git a/mooncake-store/tests/CMakeLists.txt b/mooncake-store/tests/CMakeLists.txt index 56b6485ad5..a5855d69d5 100644 --- a/mooncake-store/tests/CMakeLists.txt +++ b/mooncake-store/tests/CMakeLists.txt @@ -66,7 +66,8 @@ if(ENABLE_KV_EVENTS) endif() endif() add_store_test(master_service_test master_service_test.cpp) -add_store_test(batch_evict_test batch_evict_test.cpp) +add_ha_test(master_service_evict_scenario_test master_scenario.cpp + master_service_evict_scenario_test.cpp) add_store_test(master_scenario_test master_scenario.cpp master_scenario_test.cpp) add_store_test(master_service_scenario_test master_scenario.cpp diff --git a/mooncake-store/tests/batch_evict_test.cpp b/mooncake-store/tests/batch_evict_test.cpp deleted file mode 100644 index 84ad36a749..0000000000 --- a/mooncake-store/tests/batch_evict_test.cpp +++ /dev/null @@ -1,468 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include "master_service.h" -#include "mutex.h" -#include "types.h" -#include "utils.h" - -namespace mooncake::test { - -// Deterministic correctness tests for MasterService::BatchEvict. -// -// These tests drive BatchEvict directly instead of filling a segment and -// waiting for the background eviction thread. Lease timestamps are written -// explicitly so that the eviction order, the exact eviction count, the -// soft-pin fallback and the whole-group semantics are all observable without -// sleeps, background threads or timing assumptions. -class BatchEvictTest : public ::testing::Test { - protected: - static constexpr const char* kSegmentName = "batch_evict_test_segment"; - static constexpr size_t kSegmentBase = 0x500000000; - static constexpr size_t kSegmentSize = 256ULL * 1024 * 1024; - static constexpr uint64_t kObjectSize = 1024; - - void SetUp() override { - google::InitGoogleLogging("BatchEvictTest"); - FLAGS_logtostderr = true; - } - - void TearDown() override { google::ShutdownGoogleLogging(); } - - // Eviction ratio 0 and a 100% high watermark keep the background eviction - // thread from interfering: every eviction in these tests is the one the - // test itself requests. - static MasterServiceConfig MakeConfig(bool allow_soft_pin_eviction) { - return MasterServiceConfig::builder() - .set_memory_allocator(BufferAllocatorType::OFFSET) - .set_default_kv_lease_ttl(0) - .set_default_kv_soft_pin_ttl(60 * 60 * 1000) - .set_allow_evict_soft_pinned_objects(allow_soft_pin_eviction) - .set_eviction_ratio(0.0) - .set_eviction_high_watermark_ratio(1.0) - .set_client_live_ttl_sec(3600) - .build(); - } - - static Segment MakeSegment() { - Segment segment; - segment.id = generate_uuid(); - segment.name = kSegmentName; - segment.base = kSegmentBase; - segment.size = kSegmentSize; - segment.te_endpoint = segment.name; - return segment; - } - - static UUID MountSegment(MasterService& service) { - const UUID client_id = generate_uuid(); - auto result = service.MountSegment(MakeSegment(), client_id); - EXPECT_TRUE(result.has_value()); - return client_id; - } - - static std::string Key(size_t index) { - return "batch_evict_key_" + std::to_string(index); - } - - static void PutObject(MasterService& service, const UUID& client_id, - const std::string& key, bool with_soft_pin = false, - const std::string& group_id = std::string()) { - ReplicateConfig config; - config.replica_num = 1; - config.preferred_segment = kSegmentName; - config.soft_pin_action = - with_soft_pin ? SoftPinAction::ENABLE : SoftPinAction::PRESERVE; - if (!group_id.empty()) { - config.group_ids = std::vector{group_id}; - } - - auto put_start = service.PutStart(client_id, key, TenantId::Default(), - kObjectSize, config); - ASSERT_TRUE(put_start.has_value()) - << "PutStart failed for key=" << key - << ", error=" << toString(put_start.error()); - ASSERT_TRUE(service - .PutEnd(client_id, key, TenantId::Default(), - ReplicaType::MEMORY) - .has_value()); - } - - // Ungrouped objects are sharded by key hash, but grouped objects are - // routed by their group id, so the direct shard index is not always the - // one holding the key. Fall back to a shard scan in that case. - template - static bool WithMetadata(MasterService& service, const std::string& key, - Fn&& fn) { - const TenantId tenant = TenantId::Default(); - - auto try_shard = [&](size_t shard_idx) -> bool { - MasterService::MetadataShardAccessorRW shard(&service, shard_idx); - auto tenant_it = shard->tenants.find(tenant); - if (tenant_it == shard->tenants.end()) { - return false; - } - auto metadata_it = tenant_it->second.metadata.find(key); - if (metadata_it == tenant_it->second.metadata.end()) { - return false; - } - fn(metadata_it->second); - return true; - }; - - if (try_shard(service.getMetadataShardIndex(tenant, key))) { - return true; - } - for (size_t shard_idx = 0; shard_idx < MasterService::kNumShards; - ++shard_idx) { - if (try_shard(shard_idx)) { - return true; - } - } - ADD_FAILURE() << "metadata not found for key=" << key; - return false; - } - - static void SetLease(MasterService& service, const std::string& key, - std::chrono::system_clock::time_point lease_timeout, - std::optional - soft_pin_timeout = std::nullopt) { - WithMetadata(service, key, - [&](MasterService::ObjectMetadata& metadata) { - SpinLocker locker(&metadata.lock); - metadata.lease_timeout = lease_timeout; - metadata.soft_pin_timeout = soft_pin_timeout; - }); - } - - static bool Exists(MasterService& service, const std::string& key) { - auto result = service.ExistKey(key, TenantId::Default()); - if (!result.has_value()) { - ADD_FAILURE() << "ExistKey failed for key=" << key - << ", error=" << toString(result.error()); - return false; - } - return result.value(); - } - - static void RunBatchEvict(MasterService& service, double target, - double lowerbound) { - service.BatchEvict(target, lowerbound); - } - - static std::chrono::system_clock::time_point ExpiredBase() { - return std::chrono::system_clock::now() - std::chrono::hours(1); - } - - // Populates `count` expired objects whose lease timestamps are strictly - // increasing, so Key(i) is always older than Key(i + 1) and no timestamp - // ties exist at the eviction boundary. - static void PopulateOldestFirst(MasterService& service, - const UUID& client_id, size_t count) { - const auto base = ExpiredBase(); - for (size_t i = 0; i < count; ++i) { - PutObject(service, client_id, Key(i)); - SetLease(service, Key(i), base + std::chrono::nanoseconds(i)); - } - } - // Builds a population whose oldest `blocked_count` candidates belong to a - // group that also holds one member under an active lease. Those candidates - // pass the census — their own lease has expired — but cannot be evicted - // during execution, because a group is only evictable once every member's - // lease has expired. That is an execution-stage failure without any hook: - // it exercises the same recovery path as metadata churn between the census - // and the eviction pass. - struct BlockedGroupPopulation { - size_t blocked_count; - size_t keeper_index; - size_t plain_begin; - size_t plain_count; - size_t total_objects; - }; - - static BlockedGroupPopulation PopulateBlockedGroup( - MasterService& service, const UUID& client_id, - const std::string& group_id, size_t blocked_count, size_t plain_count) { - const auto base = ExpiredBase(); - const auto active_lease = - std::chrono::system_clock::now() + std::chrono::hours(1); - - // Oldest leases: the blocked group members are always selected first. - for (size_t i = 0; i < blocked_count; ++i) { - PutObject(service, client_id, Key(i), /*with_soft_pin=*/false, - group_id); - SetLease(service, Key(i), base + std::chrono::nanoseconds(i)); - } - - // The keeper shares the group but holds an active lease, so no member - // of the group can be evicted. - const size_t keeper_index = blocked_count; - PutObject(service, client_id, Key(keeper_index), - /*with_soft_pin=*/false, group_id); - SetLease(service, Key(keeper_index), active_lease); - - // Plain objects are strictly newer than every blocked member. - const size_t plain_begin = keeper_index + 1; - for (size_t i = 0; i < plain_count; ++i) { - const size_t index = plain_begin + i; - PutObject(service, client_id, Key(index)); - SetLease(service, Key(index), - base + std::chrono::nanoseconds(index)); - } - - return {blocked_count, keeper_index, plain_begin, plain_count, - plain_begin + plain_count}; - } - - static size_t CountAlive(MasterService& service, size_t begin, size_t end) { - size_t alive = 0; - for (size_t i = begin; i < end; ++i) { - if (Exists(service, Key(i))) { - ++alive; - } - } - return alive; - } -}; - -// Oldest-first: with distinct lease timestamps the evicted set must be exactly -// the oldest ceil(N * ratio) objects, and nothing newer. -TEST_F(BatchEvictTest, EvictsExactOldestObjectsAtLowRatio) { - constexpr size_t kObjectCount = 400; - constexpr size_t kExpectedEvicted = 20; // ceil(400 * 0.05) - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - PopulateOldestFirst(service, client_id, kObjectCount); - ASSERT_EQ(service.GetKeyCount(), kObjectCount); - - RunBatchEvict(service, /*target=*/0.05, /*lowerbound=*/0.05); - - EXPECT_EQ(service.GetKeyCount(), kObjectCount - kExpectedEvicted); - for (size_t i = 0; i < kObjectCount; ++i) { - EXPECT_EQ(Exists(service, Key(i)), i >= kExpectedEvicted) - << "unexpected eviction outcome at index=" << i; - } -} - -// target == lowerbound: the first pass already satisfies the lower bound, so -// the second pass must not evict anything extra. -TEST_F(BatchEvictTest, TargetEqualsLowerBoundEvictsExactCount) { - constexpr size_t kObjectCount = 250; - constexpr size_t kExpectedEvicted = 25; // ceil(250 * 0.10) - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - PopulateOldestFirst(service, client_id, kObjectCount); - ASSERT_EQ(service.GetKeyCount(), kObjectCount); - - RunBatchEvict(service, /*target=*/0.10, /*lowerbound=*/0.10); - - EXPECT_EQ(service.GetKeyCount(), kObjectCount - kExpectedEvicted); - EXPECT_FALSE(Exists(service, Key(kExpectedEvicted - 1))); - EXPECT_TRUE(Exists(service, Key(kExpectedEvicted))); -} - -// Soft-pin fallback: unpinned objects go first; soft-pinned objects are only -// evicted by the second pass, oldest first, and only up to the lower bound. -TEST_F(BatchEvictTest, SoftPinnedEvictedOnlyAfterUnpinned) { - constexpr size_t kNoPinCount = 10; - constexpr size_t kSoftPinCount = 10; - // ceil(20 * 0.80) == 16; the first pass can only evict the 10 unpinned - // objects, leaving 6 for the soft-pinned second pass. - constexpr size_t kExpectedSoftPinEvicted = 6; - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/true)); - const UUID client_id = MountSegment(service); - - const auto base = ExpiredBase(); - const auto active_soft_pin = - std::chrono::system_clock::now() + std::chrono::hours(1); - - for (size_t i = 0; i < kNoPinCount; ++i) { - PutObject(service, client_id, Key(i)); - SetLease(service, Key(i), base + std::chrono::nanoseconds(i)); - } - for (size_t i = 0; i < kSoftPinCount; ++i) { - const size_t index = kNoPinCount + i; - PutObject(service, client_id, Key(index), /*with_soft_pin=*/true); - SetLease(service, Key(index), base + std::chrono::nanoseconds(index), - active_soft_pin); - } - ASSERT_EQ(service.GetKeyCount(), kNoPinCount + kSoftPinCount); - - RunBatchEvict(service, /*target=*/0.80, /*lowerbound=*/0.80); - - for (size_t i = 0; i < kNoPinCount; ++i) { - EXPECT_FALSE(Exists(service, Key(i))) - << "unpinned object survived at index=" << i; - } - for (size_t i = 0; i < kSoftPinCount; ++i) { - const size_t index = kNoPinCount + i; - EXPECT_EQ(Exists(service, Key(index)), i >= kExpectedSoftPinEvicted) - << "unexpected soft-pin outcome at index=" << index; - } - EXPECT_EQ(service.GetKeyCount(), kSoftPinCount - kExpectedSoftPinEvicted); -} - -// Whole-group: selecting one group member evicts the entire group as a unit, -// even when the target only calls for a single object. -TEST_F(BatchEvictTest, WholeGroupEvictedTogether) { - constexpr size_t kObjectCount = 10; - constexpr size_t kGroupSize = 3; - const std::string group_id = "batch_evict_test_group"; - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - - const auto base = ExpiredBase(); - for (size_t i = 0; i < kObjectCount; ++i) { - PutObject(service, client_id, Key(i), /*with_soft_pin=*/false, - i < kGroupSize ? group_id : std::string()); - SetLease(service, Key(i), base + std::chrono::nanoseconds(i)); - } - - // ceil(10 * 0.10) == 1 and the oldest candidate is a group member, so the - // whole group is evicted even though only one object was requested. - RunBatchEvict(service, /*target=*/0.10, /*lowerbound=*/0.10); - - EXPECT_EQ(service.GetKeyCount(), kObjectCount - kGroupSize); - for (size_t i = 0; i < kGroupSize; ++i) { - EXPECT_FALSE(Exists(service, Key(i))) - << "group member survived at index=" << i; - } - for (size_t i = kGroupSize; i < kObjectCount; ++i) { - EXPECT_TRUE(Exists(service, Key(i))) - << "non-group object evicted at index=" << i; - } -} - -// Whole-group safety: a single group member under an active lease keeps every -// member of that group resident. -TEST_F(BatchEvictTest, UnexpiredGroupMemberBlocksWholeGroup) { - constexpr size_t kObjectCount = 10; - constexpr size_t kGroupSize = 3; - const std::string group_id = "batch_evict_test_blocked_group"; - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - - const auto base = ExpiredBase(); - const auto unexpired = - std::chrono::system_clock::now() + std::chrono::hours(1); - - for (size_t i = 0; i < kObjectCount; ++i) { - PutObject(service, client_id, Key(i), /*with_soft_pin=*/false, - i < kGroupSize ? group_id : std::string()); - SetLease(service, Key(i), base + std::chrono::nanoseconds(i)); - } - // Hold one member of the group under an active lease. - SetLease(service, Key(kGroupSize - 1), unexpired); - - RunBatchEvict(service, /*target=*/0.10, /*lowerbound=*/0.10); - - for (size_t i = 0; i < kGroupSize; ++i) { - EXPECT_TRUE(Exists(service, Key(i))) - << "blocked group member was evicted at index=" << i; - } - // The blocked group yields nothing, so exactly one ungrouped object is - // evicted instead. - EXPECT_EQ(service.GetKeyCount(), kObjectCount - 1); -} - -// High ratio: the same oldest-first and exact-count guarantees must hold when -// the requested ratio covers most of the population. -TEST_F(BatchEvictTest, HighRatioEvictsExactOldestCount) { - constexpr size_t kObjectCount = 200; - constexpr size_t kExpectedEvicted = 160; // ceil(200 * 0.80) - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - PopulateOldestFirst(service, client_id, kObjectCount); - ASSERT_EQ(service.GetKeyCount(), kObjectCount); - - RunBatchEvict(service, /*target=*/0.80, /*lowerbound=*/0.80); - - EXPECT_EQ(service.GetKeyCount(), kObjectCount - kExpectedEvicted); - for (size_t i = 0; i < kObjectCount; ++i) { - EXPECT_EQ(Exists(service, Key(i)), i >= kExpectedEvicted) - << "unexpected eviction outcome at index=" << i; - } -} - -// Reserve: a small number of candidates that pass the census but fail during -// execution is absorbed by the reserve slack, so the requested target is still -// met exactly and no refill scan is needed. -TEST_F(BatchEvictTest, ReserveAbsorbsExecutionFailuresAndMeetsTarget) { - constexpr size_t kBlockedCount = 12; - constexpr size_t kPlainCount = 1200; - // 1213 objects in total, ceil(1213 * 0.05) == 61. - constexpr size_t kExpectedEvicted = 61; - // The 61 oldest candidates are the 12 blocked members plus the 49 oldest - // plain objects, so those 49 are always among the evicted set. The - // remaining 12 evictions come from the reserve, whose internal order is - // unspecified, so only the total is asserted for them. - constexpr size_t kAlwaysEvictedPlain = 49; - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - const auto population = - PopulateBlockedGroup(service, client_id, "batch_evict_reserve_group", - kBlockedCount, kPlainCount); - ASSERT_EQ(service.GetKeyCount(), population.total_objects); - - RunBatchEvict(service, /*target=*/0.05, /*lowerbound=*/0.05); - - EXPECT_EQ(service.GetKeyCount(), - population.total_objects - kExpectedEvicted); - EXPECT_EQ(CountAlive(service, 0, kBlockedCount), kBlockedCount) - << "blocked group members must survive"; - EXPECT_TRUE(Exists(service, Key(population.keeper_index))); - for (size_t i = 0; i < kAlwaysEvictedPlain; ++i) { - EXPECT_FALSE(Exists(service, Key(population.plain_begin + i))) - << "oldest plain object survived at offset=" << i; - } - EXPECT_EQ( - CountAlive(service, population.plain_begin, population.total_objects), - kPlainCount - kExpectedEvicted); -} - -// Refill: when every candidate inside the reserve frontier fails during -// execution the reserve is exhausted, and the refill scan must recover the -// remaining objects so the requested target is still met exactly. -TEST_F(BatchEvictTest, RefillAfterReserveExhaustionStillMeetsTarget) { - // The reserve frontier spans target + max(1024, 10% of target), so more - // than 1024 blocked candidates are required to exhaust it. - constexpr size_t kBlockedCount = 1160; - constexpr size_t kPlainCount = 200; - // 1361 objects in total, ceil(1361 * 0.05) == 69. - constexpr size_t kExpectedEvicted = 69; - - MasterService service(MakeConfig(/*allow_soft_pin_eviction=*/false)); - const UUID client_id = MountSegment(service); - const auto population = - PopulateBlockedGroup(service, client_id, "batch_evict_refill_group", - kBlockedCount, kPlainCount); - ASSERT_EQ(service.GetKeyCount(), population.total_objects); - - RunBatchEvict(service, /*target=*/0.05, /*lowerbound=*/0.05); - - // Exact target attainment is the property refill exists to protect: the - // whole frontier yielded nothing, so every eviction came from the refill. - EXPECT_EQ(service.GetKeyCount(), - population.total_objects - kExpectedEvicted); - EXPECT_EQ(CountAlive(service, 0, kBlockedCount), kBlockedCount) - << "blocked group members must survive"; - EXPECT_TRUE(Exists(service, Key(population.keeper_index))); - EXPECT_EQ( - CountAlive(service, population.plain_begin, population.total_objects), - kPlainCount - kExpectedEvicted); -} - -} // namespace mooncake::test diff --git a/mooncake-store/tests/ha/master_service_ha_test.cpp b/mooncake-store/tests/ha/master_service_ha_test.cpp index a617cc2880..7e733a38b1 100644 --- a/mooncake-store/tests/ha/master_service_ha_test.cpp +++ b/mooncake-store/tests/ha/master_service_ha_test.cpp @@ -3305,94 +3305,6 @@ TEST_F(MasterServiceHATest, BatchReplicaClearSegmentReleasesAfterDurable) { EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } -TEST_F(MasterServiceHATest, BatchEvictWritesBatchRecordOpLog) { - const std::string cluster_id = "test_batch_record_batch_evict_cluster"; - auto backend = std::make_shared(); - auto service_config = MasterServiceConfig::builder() - .set_default_kv_lease_ttl(50) - .set_enable_ha(true) - .set_enable_oplog(true) - .set_cluster_id(cluster_id) - .set_oplog_batch_max_entries(1) - .build(); - MasterService service(service_config); - ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); - - auto mounted = PrepareSimpleSegment(service, "batch_evict_segment"); - OpLogBatchStorage storage(cluster_id, *backend); - OpLogBatchRecord batch; - ReadBatchEventually(storage, 1, batch); - - const std::string key = "batch_evict_key"; - PutObjectOnSegment(service, mounted.client_id, key, "batch_evict_segment"); - ReadBatchEventually(storage, 2, batch); - - std::this_thread::sleep_for(std::chrono::milliseconds(60)); - service.RunBatchEvictForTesting(/*evict_ratio_target=*/1.0, - /*evict_ratio_lowerbound=*/1.0); - ReadBatchEventually(storage, 3, batch); - - ASSERT_EQ(1u, batch.entries.size()); - EXPECT_EQ(OpType::REMOVE, batch.entries[0].op_type); - EXPECT_EQ(kDefaultTenant.value(), batch.entries[0].tenant_id); - EXPECT_EQ(key, batch.entries[0].object_key); - EXPECT_EQ(3u, batch.entries[0].sequence_id); -} - -TEST_F(MasterServiceHATest, BatchEvictReleasesMemoryAfterDurable) { - const std::string cluster_id = "test_batch_record_batch_evict_finalize"; - auto backend = std::make_shared(); - auto service_config = - MasterServiceConfig::builder() - .set_default_kv_lease_ttl(50) - .set_enable_ha(true) - .set_enable_oplog(true) - .set_cluster_id(cluster_id) - .set_oplog_batch_max_entries(1) - .set_enable_multi_tenants(true) - .set_tenant_quota_connector_type("file") - .set_tenant_quota_connector_uri( - WriteTenantPolicyFile({{kDefaultTenant.value(), 1024}})) - .build(); - MasterService service(service_config); - auto* writer = InstallGatedWriter(service, backend); - - auto mounted = PrepareSimpleSegment(service, "batch_evict_finalize_seg"); - OpLogBatchStorage storage(cluster_id, *backend); - OpLogBatchRecord batch; - ReadBatchEventually(storage, 1, batch); - - const std::string key = "batch_evict_finalize_key"; - PutObjectOnSegment(service, mounted.client_id, key, - "batch_evict_finalize_seg"); - ReadBatchEventually(storage, 2, batch); - ASSERT_TRUE(writer->PauseCallbacksAfter(batch.last_seq)); - - std::this_thread::sleep_for(std::chrono::milliseconds(60)); - backend->BlockTxn(); - service.RunBatchEvictForTesting(/*evict_ratio_target=*/1.0, - /*evict_ratio_lowerbound=*/1.0); - EXPECT_FALSE(service.GetReplicaList(key, kDefaultTenant).has_value()); - - ReplicateConfig config; - config.replica_num = 1; - const std::string before_finalize_key = "before_batch_evict_finalize_key"; - auto before_finalize = service.PutStart( - mounted.client_id, before_finalize_key, kDefaultTenant, 1024, config); - EXPECT_FALSE(before_finalize.has_value()); - - backend->AllowTxn(); - ReadBatchEventually(storage, 3, batch); - EXPECT_EQ(1024, TenantUsedBytes(service)); - ASSERT_TRUE(writer->RunCallbacksThrough(batch.last_seq)); - EXPECT_EQ(0, TenantUsedBytes(service)); - - auto after_finalize = - service.PutStart(mounted.client_id, "after_batch_evict_finalize_key", - kDefaultTenant, 1024, config); - EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); -} - TEST_F(MasterServiceHATest, EvictDiskReplicaWritesBatchRecordOpLog) { const std::string cluster_id = "test_batch_record_disk_evict_cluster"; auto backend = std::make_shared(); diff --git a/mooncake-store/tests/master_scenario.cpp b/mooncake-store/tests/master_scenario.cpp index da1032d570..9b06c867dd 100644 --- a/mooncake-store/tests/master_scenario.cpp +++ b/mooncake-store/tests/master_scenario.cpp @@ -3,8 +3,10 @@ #include #include +#include #include +#include "mutex.h" #include "types.h" namespace mooncake::test { @@ -56,10 +58,41 @@ UpsertRevokeAction UpsertRevoke(std::string key) { RemoveAction Remove(std::string key) { return {.key = std::move(key)}; } +ExpireAtAction ExpireAt(std::string key, + std::chrono::system_clock::time_point lease_timeout) { + return {.key = std::move(key), .lease_timeout = lease_timeout}; +} + +MemoryEvictAction EvictMemory(double target_ratio) { + return {.target_ratio = target_ratio, .lower_bound_ratio = target_ratio}; +} + ObjectSpec<> Object(std::string key) { return ObjectSpec<>(std::move(key)); } +ObjectsSpec<> Objects(size_t begin, size_t end) { + return ObjectsSpec<>(begin, end); +} + +ObjectsSpec<> Objects(std::initializer_list keys) { + return ObjectsSpec<>(std::vector(keys)); +} + +KeyCountSpec KeyCount(size_t value) { return {.value = value}; } + +TenantQuotaSpec TenantQuota(std::string tenant) { + return {.tenant = std::move(tenant)}; +} + +OpLogUnavailableSpec OpLogUnavailable() { return {}; } + MasterScenario::MasterScenario(std::string name) : name_(std::move(name)) {} +MasterScenario::MasterScenario(std::string name, MasterServiceConfig config, + std::shared_ptr batch_oplog_backend) + : name_(std::move(name)), + config_(std::move(config)), + batch_oplog_backend_(std::move(batch_oplog_backend)) {} + MasterScenario::~MasterScenario() = default; MasterScenario& MasterScenario::Given(MemoryNodeSpec node) { @@ -86,6 +119,53 @@ MasterScenario& MasterScenario::Given(MemoryNodeSpec node) { return *this; } +MasterScenario& MasterScenario::Given(ObjectsSpec<> objects) { + if (objects.keys.empty()) { + Fail( + "Objects requires at least one key; indexed ranges must use " + "NamedBy"); + return *this; + } + if (objects.size == 0) { + Fail("Objects requires non-zero Size"); + return *this; + } + if (objects.preferred_node.empty()) { + Fail("Objects requires CompleteOn"); + return *this; + } + + for (size_t offset = 0; offset < objects.keys.size(); ++offset) { + const auto& key = objects.keys[offset]; + auto put = PutStart(key, objects.size) + .By(objects.actor) + .ForTenant(objects.tenant) + .OnNode(objects.preferred_node); + if (!objects.group_id.empty()) { + put.InGroup(objects.group_id); + } + if (objects.with_soft_pin) { + put.WithSoftPin(); + } + if (objects.with_hard_pin) { + put.WithHardPin(); + } + When(std::move(put)); + When(PutEnd(key).By(objects.actor).ForTenant(objects.tenant)); + + if (objects.lease_timeout_base.has_value()) { + auto expire = ExpireAt(key, *objects.lease_timeout_base + + objects.lease_timeout_step * offset) + .ForTenant(objects.tenant); + if (objects.soft_pin_timeout.has_value()) { + expire.SoftPinnedUntil(*objects.soft_pin_timeout); + } + When(std::move(expire)); + } + } + return *this; +} + MasterScenario& MasterScenario::WhenPutStart(PutStartActionData action) { if (!EnsureService()) { return *this; @@ -93,9 +173,16 @@ MasterScenario& MasterScenario::WhenPutStart(PutStartActionData action) { ReplicateConfig config; config.replica_num = action.requested_replica_count; + config.preferred_segment = action.preferred_node; + config.soft_pin_action = + action.with_soft_pin ? SoftPinAction::ENABLE : SoftPinAction::PRESERVE; + config.with_hard_pin = action.with_hard_pin; + if (!action.group_id.empty()) { + config.group_ids = {action.group_id}; + } const auto result = service_->PutStart(ActorId(action.actor), action.key, - TenantId::Default(), action.size, config); + TenantId(action.tenant), action.size, config); ValidateStartResult("PutStart(" + action.key + ")", action.expected_error, action.expected_replica_count, action.expected_replica_status, result); @@ -124,8 +211,8 @@ MasterScenario& MasterScenario::When(PutEndAction action) { } const auto result = - service_->PutEnd(ActorId(action.actor), action.key, TenantId::Default(), - ReplicaType::MEMORY); + service_->PutEnd(ActorId(action.actor), action.key, + TenantId(action.tenant), ReplicaType::MEMORY); ValidateActionResult("PutEnd(" + action.key + ")", action.expected_error, result.has_value(), result ? ErrorCode::OK : result.error()); @@ -153,7 +240,7 @@ MasterScenario& MasterScenario::When(PutRevokeAction action) { const auto result = service_->PutRevoke(ActorId(action.actor), action.key, - TenantId::Default(), ReplicaType::MEMORY); + TenantId(action.tenant), ReplicaType::MEMORY); ValidateActionResult("PutRevoke(" + action.key + ")", action.expected_error, result.has_value(), result ? ErrorCode::OK : result.error()); @@ -179,13 +266,58 @@ MasterScenario& MasterScenario::When(RemoveAction action) { return *this; } - const auto result = service_->Remove(action.key, TenantId::Default()); + const auto result = service_->Remove(action.key, TenantId(action.tenant)); ValidateActionResult("Remove(" + action.key + ")", action.expected_error, result.has_value(), result ? ErrorCode::OK : result.error()); return *this; } +MasterScenario& MasterScenario::When(ExpireAtAction action) { + if (!EnsureService()) { + return *this; + } + + const TenantId tenant(action.tenant); + auto update = [&](size_t shard_idx) { + MasterService::MetadataShardAccessorRW shard(service_.get(), shard_idx); + auto tenant_it = shard->tenants.find(tenant); + if (tenant_it == shard->tenants.end()) { + return false; + } + auto metadata_it = tenant_it->second.metadata.find(action.key); + if (metadata_it == tenant_it->second.metadata.end()) { + return false; + } + SpinLocker locker(&metadata_it->second.lock); + metadata_it->second.lease_timeout = action.lease_timeout; + metadata_it->second.soft_pin_timeout = action.soft_pin_timeout; + return true; + }; + + const size_t routed = service_->getMetadataShardIndex(tenant, action.key); + if (update(routed)) { + return *this; + } + for (size_t shard_idx = 0; shard_idx < MasterService::kNumShards; + ++shard_idx) { + if (shard_idx != routed && update(shard_idx)) { + return *this; + } + } + Fail("ExpireAt(" + action.key + ") could not find object"); + return *this; +} + +MasterScenario& MasterScenario::When(MemoryEvictAction action) { + if (!EnsureService()) { + return *this; + } + service_->RunBatchEvictForTesting(action.target_ratio, + action.lower_bound_ratio); + return *this; +} + MasterScenario& MasterScenario::ThenObject(ObjectSpecData object, ObjectExpectation expectation) { if (!EnsureService()) { @@ -193,7 +325,7 @@ MasterScenario& MasterScenario::ThenObject(ObjectSpecData object, } const auto result = - service_->GetReplicaList(object.key, TenantId::Default()); + service_->GetReplicaList(object.key, TenantId(object.tenant)); if (expectation == ObjectExpectation::MISSING) { if (result) { Fail("Object(" + object.key + ") exists; expected it not to exist"); @@ -238,6 +370,98 @@ MasterScenario& MasterScenario::ThenObject(ObjectSpecData object, return *this; } +MasterScenario& MasterScenario::ThenObjects(ObjectsSpecData objects, + ObjectExpectation expectation) { + if (objects.keys.empty()) { + Fail( + "Objects requires at least one key; indexed ranges must use " + "NamedBy"); + return *this; + } + for (auto& key : objects.keys) { + ObjectSpecData object(std::move(key)); + object.tenant = objects.tenant; + ThenObject(std::move(object), expectation); + } + return *this; +} + +MasterScenario& MasterScenario::Then(KeyCountSpec key_count) { + if (!EnsureService()) { + return *this; + } + const size_t actual = service_->GetKeyCount(); + if (actual != key_count.value) { + Fail("KeyCount is " + std::to_string(actual) + "; expected " + + std::to_string(key_count.value)); + } + return *this; +} + +MasterScenario& MasterScenario::Then(TenantQuotaSpec tenant_quota) { + if (!EnsureService()) { + return *this; + } + + auto snapshot = + service_->GetTenantQuotaSnapshot(TenantId(tenant_quota.tenant)); + if (!snapshot.has_value()) { + Fail("TenantQuota(" + tenant_quota.tenant + ") is not registered"); + return *this; + } + + const auto matches = [&tenant_quota](const TenantQuotaSnapshot& value) { + return (!tenant_quota.used_bytes.has_value() || + value.used_bytes == *tenant_quota.used_bytes) && + (!tenant_quota.reserved_bytes.has_value() || + value.reserved_bytes == *tenant_quota.reserved_bytes); + }; + const auto deadline = + std::chrono::steady_clock::now() + tenant_quota.eventual_timeout; + while (!matches(*snapshot) && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + snapshot = + service_->GetTenantQuotaSnapshot(TenantId(tenant_quota.tenant)); + if (!snapshot.has_value()) { + Fail("TenantQuota(" + tenant_quota.tenant + ") is not registered"); + return *this; + } + } + + if (tenant_quota.used_bytes.has_value() && + snapshot->used_bytes != *tenant_quota.used_bytes) { + Fail("TenantQuota(" + tenant_quota.tenant + ") uses " + + std::to_string(snapshot->used_bytes) + "; expected " + + std::to_string(*tenant_quota.used_bytes)); + } + if (tenant_quota.reserved_bytes.has_value() && + snapshot->reserved_bytes != *tenant_quota.reserved_bytes) { + Fail("TenantQuota(" + tenant_quota.tenant + ") reserves " + + std::to_string(snapshot->reserved_bytes) + "; expected " + + std::to_string(*tenant_quota.reserved_bytes)); + } + return *this; +} + +MasterScenario& MasterScenario::Then(OpLogUnavailableSpec oplog) { + if (!EnsureService()) { + return *this; + } + if (!service_->ordered_oplog_writer_) { + Fail("OpLog writer is not configured"); + return *this; + } + const auto deadline = std::chrono::steady_clock::now() + oplog.timeout; + while (service_->ordered_oplog_writer_->IsAccepting() && + std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + if (service_->ordered_oplog_writer_->IsAccepting()) { + Fail("OpLog writer was expected to be unavailable"); + } + return *this; +} + bool MasterScenario::EnsureService() { if (service_) { return true; @@ -248,7 +472,16 @@ bool MasterScenario::EnsureService() { return false; } - service_ = std::make_unique(); + service_ = std::make_unique(config_); + if (batch_oplog_backend_) { + const auto result = + service_->SetBatchOpLogBackendForTesting(batch_oplog_backend_); + if (result != ErrorCode::OK) { + Fail("failed to install batch OpLog backend: " + toString(result)); + service_.reset(); + return false; + } + } for (const auto& node : nodes_) { Segment segment; segment.id = StableUuid("segment", node.name); diff --git a/mooncake-store/tests/master_scenario.h b/mooncake-store/tests/master_scenario.h index 3a1e859b5b..9614d70496 100644 --- a/mooncake-store/tests/master_scenario.h +++ b/mooncake-store/tests/master_scenario.h @@ -1,7 +1,10 @@ #pragma once +#include #include #include +#include +#include #include #include #include @@ -52,6 +55,11 @@ struct PutStartActionData { uint64_t size; std::string actor{"default"}; size_t requested_replica_count{1}; + std::string tenant{TenantId::Default().value()}; + std::string preferred_node; + std::string group_id; + bool with_soft_pin{false}; + bool with_hard_pin{false}; std::optional expected_error{}; std::optional expected_replica_count{}; std::optional expected_replica_status{}; @@ -77,6 +85,31 @@ struct PutStartAction : PutStartActionData { return *this; } + PutStartAction& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + + PutStartAction& OnNode(std::string value) { + preferred_node = std::move(value); + return *this; + } + + PutStartAction& InGroup(std::string value) { + group_id = std::move(value); + return *this; + } + + PutStartAction& WithSoftPin() { + with_soft_pin = true; + return *this; + } + + PutStartAction& WithHardPin() { + with_hard_pin = true; + return *this; + } + auto ExpectError(ErrorCode value) const requires(expectation == PutStartExpectation::UNSPECIFIED) { @@ -196,6 +229,7 @@ UpsertStartAction<> UpsertStart(std::string key, uint64_t size); struct PutEndAction { std::string key; std::string actor{"default"}; + std::string tenant{TenantId::Default().value()}; std::optional expected_error{}; PutEndAction& By(std::string value) { @@ -203,6 +237,11 @@ struct PutEndAction { return *this; } + PutEndAction& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + PutEndAction& ExpectError(ErrorCode value) { expected_error = value; return *this; @@ -232,6 +271,7 @@ UpsertEndAction UpsertEnd(std::string key); struct PutRevokeAction { std::string key; std::string actor{"default"}; + std::string tenant{TenantId::Default().value()}; std::optional expected_error{}; PutRevokeAction& By(std::string value) { @@ -239,6 +279,11 @@ struct PutRevokeAction { return *this; } + PutRevokeAction& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + PutRevokeAction& ExpectError(ErrorCode value) { expected_error = value; return *this; @@ -267,16 +312,55 @@ UpsertRevokeAction UpsertRevoke(std::string key); struct RemoveAction { std::string key; + std::string tenant{TenantId::Default().value()}; std::optional expected_error{}; RemoveAction& ExpectError(ErrorCode value) { expected_error = value; return *this; } + + RemoveAction& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } }; RemoveAction Remove(std::string key); +struct ExpireAtAction { + std::string key; + std::chrono::system_clock::time_point lease_timeout; + std::string tenant{TenantId::Default().value()}; + std::optional soft_pin_timeout{}; + + ExpireAtAction& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + + ExpireAtAction& SoftPinnedUntil( + std::chrono::system_clock::time_point value) { + soft_pin_timeout = value; + return *this; + } +}; + +ExpireAtAction ExpireAt(std::string key, + std::chrono::system_clock::time_point lease_timeout); + +struct MemoryEvictAction { + double target_ratio; + double lower_bound_ratio; + + MemoryEvictAction& ToLowerBound(double value) { + lower_bound_ratio = value; + return *this; + } +}; + +MemoryEvictAction EvictMemory(double target_ratio); + enum class ObjectExpectation { UNSPECIFIED, READABLE, @@ -294,6 +378,7 @@ struct ObjectSpecData { explicit ObjectSpecData(std::string value) : key(std::move(value)) {} std::string key; + std::string tenant{TenantId::Default().value()}; std::optional expected_replica_count{}; std::optional expected_complete_replica_count{}; @@ -327,6 +412,11 @@ struct ObjectSpec : ObjectSpecData { return ObjectSpec(*this); } + ObjectSpec& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + auto HasReplicas(size_t value) const requires(expectation != ObjectExpectation::NOT_READY && expectation != ObjectExpectation::MISSING) @@ -355,15 +445,201 @@ struct ObjectSpec : ObjectSpecData { ObjectSpec<> Object(std::string key); +struct ObjectsSpecData { + std::vector indices; + std::vector keys; + uint64_t size{0}; + std::string actor{"default"}; + std::string tenant{TenantId::Default().value()}; + std::string preferred_node; + std::string group_id; + bool with_soft_pin{false}; + bool with_hard_pin{false}; + std::optional lease_timeout_base{}; + std::chrono::nanoseconds lease_timeout_step{1}; + std::optional soft_pin_timeout{}; +}; + +template +struct ObjectsSpec : ObjectsSpecData { + ObjectsSpec(size_t begin, size_t end) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + indices.reserve(end > begin ? end - begin : 0); + for (size_t index = begin; index < end; ++index) { + indices.push_back(index); + } + } + + explicit ObjectsSpec(std::vector values) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + keys = std::move(values); + } + + template + ObjectsSpec& NamedBy(KeyFactory&& key_factory) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + keys.clear(); + keys.reserve(indices.size()); + for (const size_t index : indices) { + keys.push_back(std::invoke(key_factory, index)); + } + return *this; + } + + ObjectsSpec& Size(uint64_t value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + size = value; + return *this; + } + + ObjectsSpec& By(std::string value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + actor = std::move(value); + return *this; + } + + ObjectsSpec& ForTenant(std::string value) { + tenant = std::move(value); + return *this; + } + + ObjectsSpec& CompleteOn(std::string value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + preferred_node = std::move(value); + return *this; + } + + ObjectsSpec& InGroup(std::string value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + group_id = std::move(value); + return *this; + } + + ObjectsSpec& WithSoftPin() + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + with_soft_pin = true; + return *this; + } + + ObjectsSpec& WithHardPin() + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + with_hard_pin = true; + return *this; + } + + ObjectsSpec& ExpiredFrom( + std::chrono::system_clock::time_point value, + std::chrono::nanoseconds step = std::chrono::nanoseconds(1)) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + lease_timeout_base = value; + lease_timeout_step = step; + return *this; + } + + ObjectsSpec& ExpiresAt(std::chrono::system_clock::time_point value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + lease_timeout_base = value; + lease_timeout_step = std::chrono::nanoseconds::zero(); + return *this; + } + + ObjectsSpec& SoftPinnedUntil(std::chrono::system_clock::time_point value) + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + with_soft_pin = true; + soft_pin_timeout = value; + return *this; + } + + auto AreReadable() const + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + return ObjectsSpec(*this); + } + + auto AreNotReady() const + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + return ObjectsSpec(*this); + } + + auto DoNotExist() const + requires(expectation == ObjectExpectation::UNSPECIFIED) + { + return ObjectsSpec(*this); + } + + private: + template + friend struct ObjectsSpec; + + template + ObjectsSpec(const ObjectsSpec& objects) : ObjectsSpecData(objects) {} +}; + +ObjectsSpec<> Objects(size_t begin, size_t end); +ObjectsSpec<> Objects(std::initializer_list keys); + +struct KeyCountSpec { + size_t value; +}; + +KeyCountSpec KeyCount(size_t value); + +struct TenantQuotaSpec { + std::string tenant; + std::optional used_bytes{}; + std::optional reserved_bytes{}; + std::chrono::milliseconds eventual_timeout{}; + + TenantQuotaSpec& Uses(uint64_t value) { + used_bytes = value; + return *this; + } + + TenantQuotaSpec& Reserves(uint64_t value) { + reserved_bytes = value; + return *this; + } + + TenantQuotaSpec& Eventually( + std::chrono::milliseconds timeout = std::chrono::seconds(1)) { + eventual_timeout = timeout; + return *this; + } +}; + +TenantQuotaSpec TenantQuota(std::string tenant); + +struct OpLogUnavailableSpec { + std::chrono::milliseconds timeout{std::chrono::seconds(1)}; +}; + +OpLogUnavailableSpec OpLogUnavailable(); + class MasterScenario { public: explicit MasterScenario(std::string name); + MasterScenario(std::string name, MasterServiceConfig config, + std::shared_ptr batch_oplog_backend = nullptr); ~MasterScenario(); MasterScenario(const MasterScenario&) = delete; MasterScenario& operator=(const MasterScenario&) = delete; MasterScenario& Given(MemoryNodeSpec node); + MasterScenario& Given(ObjectsSpec<> objects); template MasterScenario& When(PutStartAction action) { return WhenPutStart(std::move(action)); @@ -378,12 +654,22 @@ class MasterScenario { MasterScenario& When(PutRevokeAction action); MasterScenario& When(UpsertRevokeAction action); MasterScenario& When(RemoveAction action); + MasterScenario& When(ExpireAtAction action); + MasterScenario& When(MemoryEvictAction action); template requires(expectation != ObjectExpectation::UNSPECIFIED) MasterScenario& Then(ObjectSpec object) { return ThenObject(std::move(object), expectation); } + template + requires(expectation != ObjectExpectation::UNSPECIFIED) + MasterScenario& Then(ObjectsSpec objects) { + return ThenObjects(std::move(objects), expectation); + } + MasterScenario& Then(KeyCountSpec key_count); + MasterScenario& Then(TenantQuotaSpec tenant_quota); + MasterScenario& Then(OpLogUnavailableSpec oplog); private: using StartResult = @@ -393,6 +679,8 @@ class MasterScenario { MasterScenario& WhenUpsertStart(UpsertStartActionData action); MasterScenario& ThenObject(ObjectSpecData object, ObjectExpectation expectation); + MasterScenario& ThenObjects(ObjectsSpecData objects, + ObjectExpectation expectation); bool EnsureService(); UUID ActorId(std::string_view actor); void ValidateActionResult(std::string_view action, @@ -406,6 +694,8 @@ class MasterScenario { void Fail(std::string message) const; std::string name_; + MasterServiceConfig config_; + std::shared_ptr batch_oplog_backend_; bool declarations_frozen_{false}; uintptr_t next_segment_base_{0x300000000}; std::vector nodes_; diff --git a/mooncake-store/tests/master_scenario_test.cpp b/mooncake-store/tests/master_scenario_test.cpp index 0edb9034ad..7015a120f7 100644 --- a/mooncake-store/tests/master_scenario_test.cpp +++ b/mooncake-store/tests/master_scenario_test.cpp @@ -23,6 +23,9 @@ concept SupportsIsReadable = requires(T value) { value.IsReadable(); }; template concept SupportsIsNotReady = requires(T value) { value.IsNotReady(); }; +template +concept SupportsDoesNotExist = requires(T value) { value.DoesNotExist(); }; + template concept SupportsHasReplicas = requires(T value) { value.HasReplicas(1); }; @@ -30,9 +33,6 @@ template concept SupportsHasCompleteReplicas = requires(T value) { value.HasCompleteReplicas(1); }; -template -concept SupportsDoesNotExist = requires(T value) { value.DoesNotExist(); }; - template concept SupportsExpectedErrorMutation = requires(T value) { value.expected_error = ErrorCode::INTERNAL_ERROR; }; @@ -61,8 +61,11 @@ using SuccessExpectedUpsertStart = decltype(UpsertStart("compile-time", 1_KB).ExpectReplicas(1)); using UnspecifiedObject = decltype(Object("compile-time")); using NotReadyObject = decltype(Object("compile-time").IsNotReady()); -using ReadableObject = decltype(Object("compile-time").HasReplicas(1)); using MissingObject = decltype(Object("compile-time").DoesNotExist()); +using ReadableObject = decltype(Object("compile-time").HasReplicas(1)); +using UnspecifiedObjects = decltype(Objects(0, 1)); +using MissingObjects = decltype(Objects(0, 1).DoNotExist()); +using ReadableObjects = decltype(Objects(0, 1).AreReadable()); static_assert(!SupportsExpectReplicas); static_assert(!SupportsExpectStatus); @@ -87,8 +90,11 @@ static_assert(!SupportsExpectedReplicaCountMutation); static_assert(!SupportsExpectedCompleteReplicaCountMutation); static_assert(!SupportsThen); static_assert(SupportsThen); -static_assert(SupportsThen); static_assert(SupportsThen); +static_assert(SupportsThen); +static_assert(!SupportsThen); +static_assert(SupportsThen); +static_assert(SupportsThen); } // namespace @@ -225,6 +231,14 @@ TEST(MasterScenarioContractTest, ReportsNotReadyObjectWhenAbsenceExpected) { "OBJECT_NOT_FOUND"); } +TEST(MasterScenarioContractTest, ReportsKeyCountMismatch) { + EXPECT_NONFATAL_FAILURE(MasterScenario("key count mismatch") + .Given(MemoryNode("memory")) + .When(PutStart("key", 1_KB)) + .Then(KeyCount(0)), + "KeyCount is 1; expected 0"); +} + TEST(MasterScenarioContractTest, ReportsObjectReplicaCountMismatch) { EXPECT_NONFATAL_FAILURE(MasterScenario("object replica count mismatch") .Given(MemoryNode("memory")) @@ -243,4 +257,30 @@ TEST(MasterScenarioContractTest, ReportsCompleteReplicaCountMismatch) { "Object(key) has 1 complete replicas; expected 0"); } +TEST(MasterScenarioContractTest, CreatesAndChecksObjectCollections) { + MasterScenario("object collections") + .Given(MemoryNode("memory")) + .Given(Objects(2, 5) + .NamedBy([](size_t index) { + return "collection-" + std::to_string(index); + }) + .Size(1_KB) + .CompleteOn("memory")) + .Then(Objects(2, 5) + .NamedBy([](size_t index) { + return "collection-" + std::to_string(index); + }) + .AreReadable()) + .Then(KeyCount(3)); +} + +TEST(MasterScenarioContractTest, CollectionFailureIdentifiesObjectKey) { + EXPECT_NONFATAL_FAILURE( + MasterScenario("collection failure") + .Given(MemoryNode("memory")) + .Given(Objects({"present"}).Size(1_KB).CompleteOn("memory")) + .Then(Objects({"present", "missing"}).AreReadable()), + "Object(missing) is not readable: OBJECT_NOT_FOUND"); +} + } // namespace mooncake::test diff --git a/mooncake-store/tests/master_service_evict_scenario_test.cpp b/mooncake-store/tests/master_service_evict_scenario_test.cpp new file mode 100644 index 0000000000..bbc2e29095 --- /dev/null +++ b/mooncake-store/tests/master_service_evict_scenario_test.cpp @@ -0,0 +1,592 @@ +#include "master_scenario.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "ha/kv/ha_kv_backend.h" +#include "ha/oplog/oplog_batch_storage.h" +#include "ha/oplog/oplog_batch_types.h" +#include "tenant_quota_policy_store.h" +#include "types.h" + +namespace mooncake::test { +namespace { + +class EvictFakeBatchHaKvBackend : public HaKvBackend { + public: + ErrorCode Get(std::string_view key, std::string& value) override { + std::lock_guard lock(kvs_mutex_); + auto it = kvs_.find(std::string(key)); + if (it == kvs_.end()) { + return ErrorCode::ETCD_KEY_NOT_EXIST; + } + value = it->second; + return ErrorCode::OK; + } + + ErrorCode Put(std::string_view key, std::string_view value) override { + std::lock_guard lock(kvs_mutex_); + kvs_[std::string(key)] = std::string(value); + return ErrorCode::OK; + } + + ErrorCode Range(std::string_view begin_key, std::string_view end_key, + size_t limit, std::vector& kvs) override { + std::lock_guard lock(kvs_mutex_); + kvs.clear(); + for (auto it = kvs_.lower_bound(std::string(begin_key)); + it != kvs_.end() && it->first < end_key; ++it) { + kvs.push_back({.key = it->first, .value = it->second}); + if (limit != 0 && kvs.size() >= limit) { + break; + } + } + return ErrorCode::OK; + } + + bool SupportsTxn() const override { return true; } + + ErrorCode Txn(const KvTxn& txn) override { + std::lock_guard lock(kvs_mutex_); + for (const auto& compare : txn.compares) { + auto it = kvs_.find(compare.key); + if (compare.kind == KvCompareKind::kKeyNotExists) { + if (it != kvs_.end()) { + return ErrorCode::ETCD_TRANSACTION_FAIL; + } + } else if (it == kvs_.end() || + it->second != compare.expected_value) { + return ErrorCode::ETCD_TRANSACTION_FAIL; + } + } + for (const auto& put : txn.puts) { + kvs_[put.key] = put.value; + } + return ErrorCode::OK; + } + + private: + std::mutex kvs_mutex_; + std::map kvs_; +}; + +class EvictBlockingBatchHaKvBackend : public EvictFakeBatchHaKvBackend { + public: + void BlockTxn() { + std::lock_guard lock(block_mutex_); + blocked_ = true; + } + + void AllowTxn() { + { + std::lock_guard lock(block_mutex_); + blocked_ = false; + } + block_cv_.notify_all(); + } + + ErrorCode Txn(const KvTxn& txn) override { + { + std::unique_lock lock(block_mutex_); + block_cv_.wait(lock, [this] { return !blocked_; }); + } + return EvictFakeBatchHaKvBackend::Txn(txn); + } + + private: + std::mutex block_mutex_; + std::condition_variable block_cv_; + bool blocked_{false}; +}; + +class EvictFailingBatchHaKvBackend : public EvictFakeBatchHaKvBackend { + public: + void FailTransactionsWith(ErrorCode error) { + std::lock_guard lock(failure_mutex_); + transaction_error_ = error; + transaction_calls_ = 0; + } + + bool WaitForTransactionCalls( + size_t count, + std::chrono::milliseconds timeout = std::chrono::seconds(1)) { + std::unique_lock lock(failure_mutex_); + return failure_cv_.wait_for( + lock, timeout, [&] { return transaction_calls_ >= count; }); + } + + ErrorCode Txn(const KvTxn& txn) override { + ErrorCode error; + { + std::lock_guard lock(failure_mutex_); + ++transaction_calls_; + error = transaction_error_; + } + failure_cv_.notify_all(); + if (error != ErrorCode::OK) { + return error; + } + return EvictFakeBatchHaKvBackend::Txn(txn); + } + + private: + std::mutex failure_mutex_; + std::condition_variable failure_cv_; + ErrorCode transaction_error_{ErrorCode::OK}; + size_t transaction_calls_{0}; +}; + +class MasterServiceEvictScenarioTest : public ::testing::Test { + protected: + static constexpr uint64_t kObjectSize = 1_KB; + + static void SetUpTestSuite() { + google::InitGoogleLogging("MasterServiceEvictScenarioTest"); + FLAGS_logtostderr = true; + } + + static void TearDownTestSuite() { google::ShutdownGoogleLogging(); } + + void TearDown() override { + for (const auto& path : policy_files_) { + std::error_code error; + std::filesystem::remove(path, error); + } + } + + MasterServiceConfig EvictConfig(bool allow_soft_pin_eviction = false) { + return MasterServiceConfig::builder() + .set_memory_allocator(BufferAllocatorType::OFFSET) + .set_default_kv_lease_ttl(0) + .set_default_kv_soft_pin_ttl(60 * 60 * 1000) + .set_allow_evict_soft_pinned_objects(allow_soft_pin_eviction) + .set_eviction_ratio(0.0) + .set_eviction_high_watermark_ratio(1.0) + .set_client_live_ttl_sec(3600) + .build(); + } + + MasterServiceConfig TenantConfig( + const std::map& quotas) { + auto config = EvictConfig(); + config.enable_multi_tenants = true; + config.tenant_quota_connector_type = "file"; + config.tenant_quota_connector_uri = WritePolicyFile(quotas); + return config; + } + + MasterServiceConfig HaConfig( + const std::string& cluster_id, + const std::map& quotas = {}) { + auto config = EvictConfig(); + config.enable_ha = true; + config.enable_oplog = true; + config.cluster_id = cluster_id; + config.oplog_batch_max_entries = 1; + if (!quotas.empty()) { + config.enable_multi_tenants = true; + config.tenant_quota_connector_type = "file"; + config.tenant_quota_connector_uri = WritePolicyFile(quotas); + } + return config; + } + + static std::string Key(size_t index) { + return "evict_scenario_key_" + std::to_string(index); + } + + static std::chrono::system_clock::time_point ExpiredBase() { + return std::chrono::system_clock::now() - std::chrono::hours(1); + } + + static ObjectsSpec<> IndexedObjects(size_t begin, size_t end) { + auto objects = Objects(begin, end); + objects.NamedBy(Key); + return objects; + } + + void ReadBatchEventually(OpLogBatchStorage& storage, uint64_t batch_id, + OpLogBatchRecord& batch) { + ErrorCode error = ErrorCode::ETCD_KEY_NOT_EXIST; + for (int attempt = 0; attempt < 100; ++attempt) { + error = storage.ReadBatch(batch_id, batch); + if (error == ErrorCode::OK) { + break; + } + std::this_thread::sleep_for(std::chrono::milliseconds(10)); + } + ASSERT_EQ(error, ErrorCode::OK); + } + + private: + std::string WritePolicyFile(const std::map& quotas) { + TenantQuotaPolicySnapshot snapshot; + snapshot.tenant_quotas = quotas; + const auto path = + std::filesystem::temp_directory_path() / + ("mooncake_evict_scenario_quota_" + std::to_string(::getpid()) + + "_" + std::to_string(policy_files_.size()) + ".yaml"); + std::ofstream output(path); + output << FormatTenantQuotaPolicyYaml(snapshot); + output.close(); + policy_files_.push_back(path.string()); + return path.string(); + } + + std::vector policy_files_; +}; + +TEST_F(MasterServiceEvictScenarioTest, EvictsExactOldestObjectsAtLowRatio) { + constexpr size_t kObjectCount = 400; + constexpr size_t kExpectedEvicted = 20; + + MasterScenario scenario("evict exact oldest objects", EvictConfig()); + scenario.Given(MemoryNode("memory").Capacity(256 * 1024 * 1024)) + .Given(IndexedObjects(0, kObjectCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(ExpiredBase())) + .When(EvictMemory(0.05)) + .Then(KeyCount(kObjectCount - kExpectedEvicted)) + .Then(IndexedObjects(0, kExpectedEvicted).DoNotExist()) + .Then(IndexedObjects(kExpectedEvicted, kObjectCount).AreReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, TargetEqualsLowerBoundEvictsExactCount) { + constexpr size_t kObjectCount = 250; + constexpr size_t kExpectedEvicted = 25; + + MasterScenario scenario("equal target and lower bound", EvictConfig()); + scenario.Given(MemoryNode("memory").Capacity(256 * 1024 * 1024)) + .Given(IndexedObjects(0, kObjectCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(ExpiredBase())) + .When(EvictMemory(0.10).ToLowerBound(0.10)) + .Then(KeyCount(kObjectCount - kExpectedEvicted)) + .Then(Object(Key(kExpectedEvicted - 1)).DoesNotExist()) + .Then(Object(Key(kExpectedEvicted)).IsReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, SoftPinnedObjectsAreFallbackCandidates) { + constexpr size_t kUnpinnedCount = 10; + constexpr size_t kSoftPinnedCount = 10; + constexpr size_t kExpectedSoftPinnedEvicted = 6; + const auto base = ExpiredBase(); + const auto active_soft_pin = + std::chrono::system_clock::now() + std::chrono::hours(1); + + MasterScenario scenario("soft pin fallback", EvictConfig(true)); + scenario.Given(MemoryNode("memory")) + .Given(IndexedObjects(0, kUnpinnedCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base)) + .Given(IndexedObjects(kUnpinnedCount, kUnpinnedCount + kSoftPinnedCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base + std::chrono::nanoseconds(kUnpinnedCount)) + .SoftPinnedUntil(active_soft_pin)) + .When(EvictMemory(0.80)) + .Then(KeyCount(kSoftPinnedCount - kExpectedSoftPinnedEvicted)) + .Then(IndexedObjects(0, kUnpinnedCount + kExpectedSoftPinnedEvicted) + .DoNotExist()) + .Then(IndexedObjects(kUnpinnedCount + kExpectedSoftPinnedEvicted, + kUnpinnedCount + kSoftPinnedCount) + .AreReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, EvictsWholeGroupTogether) { + constexpr size_t kObjectCount = 10; + constexpr size_t kGroupSize = 3; + const auto base = ExpiredBase(); + + MasterScenario scenario("whole group eviction", EvictConfig()); + scenario.Given(MemoryNode("memory")) + .Given(IndexedObjects(0, kGroupSize) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("group") + .ExpiredFrom(base)) + .Given(IndexedObjects(kGroupSize, kObjectCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base + std::chrono::nanoseconds(kGroupSize))) + .When(EvictMemory(0.10)) + .Then(KeyCount(kObjectCount - kGroupSize)) + .Then(IndexedObjects(0, kGroupSize).DoNotExist()) + .Then(IndexedObjects(kGroupSize, kObjectCount).AreReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, ActiveGroupMemberBlocksWholeGroup) { + constexpr size_t kObjectCount = 10; + constexpr size_t kGroupSize = 3; + const auto base = ExpiredBase(); + + MasterScenario scenario("active group member blocks eviction", + EvictConfig()); + scenario.Given(MemoryNode("memory")) + .Given(IndexedObjects(0, kGroupSize) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("group") + .ExpiredFrom(base)) + .Given(IndexedObjects(kGroupSize, kObjectCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base + std::chrono::nanoseconds(kGroupSize))) + .When(ExpireAt(Key(kGroupSize - 1), std::chrono::system_clock::now() + + std::chrono::hours(1))) + .When(EvictMemory(0.10)) + .Then(KeyCount(kObjectCount - 1)) + .Then(IndexedObjects(0, kGroupSize).AreReadable()) + .Then(Object(Key(kGroupSize)).DoesNotExist()); +} + +TEST_F(MasterServiceEvictScenarioTest, EvictsExactOldestObjectsAtHighRatio) { + constexpr size_t kObjectCount = 200; + constexpr size_t kExpectedEvicted = 160; + + MasterScenario scenario("high ratio eviction", EvictConfig()); + scenario.Given(MemoryNode("memory").Capacity(256 * 1024 * 1024)) + .Given(IndexedObjects(0, kObjectCount) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(ExpiredBase())) + .When(EvictMemory(0.80)) + .Then(KeyCount(kObjectCount - kExpectedEvicted)) + .Then(IndexedObjects(0, kExpectedEvicted).DoNotExist()) + .Then(IndexedObjects(kExpectedEvicted, kObjectCount).AreReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, + ReserveAbsorbsBlockedCandidatesAndMeetsTarget) { + constexpr size_t kBlockedCount = 12; + constexpr size_t kPlainCount = 1200; + constexpr size_t kExpectedEvicted = 61; + constexpr size_t kAlwaysEvictedPlain = 49; + const auto base = ExpiredBase(); + + MasterScenario scenario("reserve absorbs blocked candidates", + EvictConfig()); + const size_t keeper = kBlockedCount; + const size_t plain_begin = keeper + 1; + const size_t total = plain_begin + kPlainCount; + + scenario.Given(MemoryNode("memory").Capacity(256 * 1024 * 1024)) + .Given(IndexedObjects(0, kBlockedCount) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("blocked-group") + .ExpiredFrom(base)) + .Given(IndexedObjects(keeper, keeper + 1) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("blocked-group") + .ExpiresAt(std::chrono::system_clock::now() + + std::chrono::hours(1))) + .Given(IndexedObjects(plain_begin, total) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base + std::chrono::nanoseconds(plain_begin))) + .When(EvictMemory(0.05)) + .Then(KeyCount(total - kExpectedEvicted)) + .Then(IndexedObjects(0, kBlockedCount + 1).AreReadable()) + .Then(IndexedObjects(plain_begin, plain_begin + kAlwaysEvictedPlain) + .DoNotExist()); +} + +TEST_F(MasterServiceEvictScenarioTest, + RefillAfterReserveExhaustionStillMeetsTarget) { + constexpr size_t kBlockedCount = 1160; + constexpr size_t kPlainCount = 200; + constexpr size_t kExpectedEvicted = 69; + const auto base = ExpiredBase(); + + MasterScenario scenario("refill after reserve exhaustion", EvictConfig()); + const size_t keeper = kBlockedCount; + const size_t plain_begin = keeper + 1; + const size_t total = plain_begin + kPlainCount; + + scenario.Given(MemoryNode("memory").Capacity(256 * 1024 * 1024)) + .Given(IndexedObjects(0, kBlockedCount) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("blocked-group") + .ExpiredFrom(base)) + .Given(IndexedObjects(keeper, keeper + 1) + .Size(kObjectSize) + .CompleteOn("memory") + .InGroup("blocked-group") + .ExpiresAt(std::chrono::system_clock::now() + + std::chrono::hours(1))) + .Given(IndexedObjects(plain_begin, total) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiredFrom(base + std::chrono::nanoseconds(plain_begin))) + .When(EvictMemory(0.05)) + .Then(KeyCount(total - kExpectedEvicted)) + .Then(IndexedObjects(0, kBlockedCount + 1).AreReadable()); +} + +TEST_F(MasterServiceEvictScenarioTest, OpLogRecordsEvictedTenantAndKey) { + const std::string cluster_id = "evict_scenario_oplog"; + auto backend = std::make_shared(); + MasterScenario scenario("eviction writes tenant-scoped oplog", + HaConfig(cluster_id, {{"tenant-a", kObjectSize}}), + backend); + scenario.Given(MemoryNode("memory")) + .Given(Objects({"cold"}) + .Size(kObjectSize) + .ForTenant("tenant-a") + .CompleteOn("memory") + .ExpiresAt(ExpiredBase())); + + OpLogBatchStorage storage(cluster_id, *backend); + OpLogBatchRecord batch; + ReadBatchEventually(storage, 2, batch); + + scenario.When(EvictMemory(1.0)); + ReadBatchEventually(storage, 3, batch); + + ASSERT_EQ(batch.entries.size(), 1); + EXPECT_EQ(batch.entries[0].op_type, OpType::REMOVE); + EXPECT_EQ(batch.entries[0].tenant_id, "tenant-a"); + EXPECT_EQ(batch.entries[0].object_key, "cold"); +} + +TEST_F(MasterServiceEvictScenarioTest, + OpLogDurabilityGatesTenantQuotaReclamation) { + const std::string cluster_id = "evict_scenario_durability"; + const std::string tenant = TenantId::Default().value(); + auto backend = std::make_shared(); + MasterScenario scenario("durability gates eviction quota reclamation", + HaConfig(cluster_id, {{tenant, kObjectSize}}), + backend); + scenario.Given(MemoryNode("memory").Capacity(kObjectSize)) + .Given(Objects({"cold"}) + .Size(kObjectSize) + .ForTenant(tenant) + .CompleteOn("memory") + .ExpiresAt(ExpiredBase())); + + OpLogBatchStorage storage(cluster_id, *backend); + OpLogBatchRecord batch; + ReadBatchEventually(storage, 2, batch); + + backend->BlockTxn(); + scenario.When(EvictMemory(1.0)) + .Then(Object("cold").DoesNotExist()) + .Then(TenantQuota(tenant).Uses(kObjectSize).Reserves(0)) + .When(PutStart("before-durable", kObjectSize) + .ForTenant(tenant) + .ExpectError(ErrorCode::TENANT_QUOTA_EXCEEDED)); + + backend->AllowTxn(); + ReadBatchEventually(storage, 3, batch); + scenario.Then(TenantQuota(tenant).Uses(0).Reserves(0).Eventually()) + .When(PutStart("after-durable", kObjectSize) + .ForTenant(tenant) + .ExpectReplicas(1)) + .Then(TenantQuota(tenant).Uses(0).Reserves(kObjectSize)); +} + +TEST_F(MasterServiceEvictScenarioTest, + OpLogReservationFailureLeavesEvictionCandidateReadable) { + const std::string cluster_id = "evict_scenario_oplog_failure"; + auto backend = std::make_shared(); + MasterScenario scenario("oplog failure keeps eviction candidate intact", + HaConfig(cluster_id), backend); + scenario.Given(MemoryNode("memory")) + .Given(Objects({"cold"}) + .Size(kObjectSize) + .CompleteOn("memory") + .ExpiresAt(ExpiredBase())); + + OpLogBatchStorage storage(cluster_id, *backend); + OpLogBatchRecord batch; + ReadBatchEventually(storage, 2, batch); + + // Fail a preceding hard-pinned object's PUT_END to put the ordered writer + // into its terminal failure state. The following eviction cannot reserve + // an OpLog slot and therefore must not mutate the cold object. + backend->FailTransactionsWith(ErrorCode::INTERNAL_ERROR); + scenario + .When(PutStart("writer-failure", kObjectSize) + .OnNode("memory") + .WithHardPin()) + .When(PutEnd("writer-failure")); + ASSERT_TRUE(backend->WaitForTransactionCalls(1)); + + scenario.Then(OpLogUnavailable()) + .When(EvictMemory(1.0)) + .Then(Object("cold").IsReadable()) + .Then(KeyCount(2)); +} + +TEST_F(MasterServiceEvictScenarioTest, + GlobalEvictReclaimsOnlySelectedTenantQuota) { + MasterScenario scenario( + "global eviction preserves tenant isolation", + TenantConfig({{"tenant-a", kObjectSize}, {"tenant-b", kObjectSize}})); + const auto base = ExpiredBase(); + scenario.Given(MemoryNode("memory")) + .Given(Objects({"same-key"}) + .Size(kObjectSize) + .ForTenant("tenant-a") + .CompleteOn("memory") + .ExpiresAt(base)) + .Given(Objects({"same-key"}) + .Size(kObjectSize) + .ForTenant("tenant-b") + .CompleteOn("memory") + .ExpiresAt(base + std::chrono::seconds(1))) + .When(EvictMemory(0.5)) + .Then(Object("same-key").ForTenant("tenant-a").DoesNotExist()) + .Then(Object("same-key").ForTenant("tenant-b").IsReadable()) + .Then(TenantQuota("tenant-a").Uses(0).Reserves(0)) + .Then(TenantQuota("tenant-b").Uses(kObjectSize).Reserves(0)); +} + +TEST_F(MasterServiceEvictScenarioTest, + TenantAdmissionEvictsOnlyThatTenantsExpiredObject) { + MasterScenario scenario( + "tenant admission evicts within tenant", + TenantConfig({{"tenant-a", kObjectSize}, {"tenant-b", kObjectSize}})); + scenario.Given(MemoryNode("memory")) + .Given(Objects({"tenant-a-old"}) + .Size(kObjectSize) + .ForTenant("tenant-a") + .CompleteOn("memory") + .ExpiresAt(ExpiredBase())) + .Given(Objects({"tenant-b-object"}) + .Size(kObjectSize) + .ForTenant("tenant-b") + .CompleteOn("memory")) + .When(PutStart("tenant-a-new", kObjectSize) + .ForTenant("tenant-a") + .ExpectReplicas(1)) + .Then(Object("tenant-a-old").ForTenant("tenant-a").DoesNotExist()) + .Then(Object("tenant-b-object").ForTenant("tenant-b").IsReadable()) + .Then(TenantQuota("tenant-a").Uses(0).Reserves(kObjectSize)) + .When(PutEnd("tenant-a-new").ForTenant("tenant-a")) + .Then(TenantQuota("tenant-a").Uses(kObjectSize).Reserves(0)) + .Then(TenantQuota("tenant-b").Uses(kObjectSize).Reserves(0)); +} + +} // namespace +} // namespace mooncake::test From 1597268732043f92823424e053d057222d14274d Mon Sep 17 00:00:00 2001 From: LZW <99333079+Lin-z-w@users.noreply.github.com> Date: Wed, 12 Aug 2026 17:49:00 +0800 Subject: [PATCH 044/483] [Store] Make tenant quota charge and release lock-free (#3162) Co-authored-by: yokinoshitayoki --- .../mooncake-store-deployment-guide.md | 6 +- docs/source/deployment/multi-tenancy.md | 11 +- mooncake-store/include/master_service.h | 83 +- mooncake-store/include/tenant_quota.h | 109 +- mooncake-store/include/tenant_quota_ledger.h | 65 ++ mooncake-store/include/tenant_quota_sharded.h | 27 +- .../include/tenant_quota_sharded_impl.h | 116 +-- mooncake-store/src/CMakeLists.txt | 1 + mooncake-store/src/master_admin_service.cpp | 54 +- mooncake-store/src/master_service.cpp | 942 ++++++++++-------- mooncake-store/src/tenant_quota.cpp | 525 +++++----- mooncake-store/src/tenant_quota_ledger.cpp | 196 ++++ .../src/tenant_quota_policy_store.cpp | 9 +- mooncake-store/tests/CMakeLists.txt | 1 + .../tests/ha/master_service_ha_test.cpp | 113 ++- .../tests/master_admin_server_test.cpp | 23 +- .../master_service_tenant_quota_test.cpp | 417 +++++++- .../tests/promotion_on_hit_test.cpp | 92 ++ .../tests/tenant_quota_ledger_test.cpp | 219 ++++ mooncake-store/tests/tenant_quota_test.cpp | 238 +++-- 20 files changed, 2217 insertions(+), 1030 deletions(-) create mode 100644 mooncake-store/include/tenant_quota_ledger.h create mode 100644 mooncake-store/src/tenant_quota_ledger.cpp create mode 100644 mooncake-store/tests/tenant_quota_ledger_test.cpp diff --git a/docs/source/deployment/mooncake-store-deployment-guide.md b/docs/source/deployment/mooncake-store-deployment-guide.md index 7ac43272dc..a269628a96 100644 --- a/docs/source/deployment/mooncake-store-deployment-guide.md +++ b/docs/source/deployment/mooncake-store-deployment-guide.md @@ -386,10 +386,8 @@ When tenant quota is enabled, `/metrics` also includes per-tenant quota gauges a - `mooncake_tenant_quota_requested_bytes{tenant_id}` - `mooncake_tenant_quota_effective_bytes{tenant_id}` -- `mooncake_tenant_quota_used_bytes{tenant_id}` -- `mooncake_tenant_quota_reserved_bytes{tenant_id}` -- `mooncake_tenant_quota_committed_count{tenant_id}` -- `mooncake_tenant_quota_metadata_object_count{tenant_id}` +- `mooncake_tenant_quota_charged_bytes{tenant_id}` +- `mooncake_tenant_quota_admission_closed{tenant_id}` - `mooncake_tenant_quota_over_quota{tenant_id}` - `mooncake_tenant_quota_explicit_policy{tenant_id}` - `mooncake_tenant_quota_reject_total{tenant_id,reason}` diff --git a/docs/source/deployment/multi-tenancy.md b/docs/source/deployment/multi-tenancy.md index b4c48cf306..8c8614eed4 100644 --- a/docs/source/deployment/multi-tenancy.md +++ b/docs/source/deployment/multi-tenancy.md @@ -55,7 +55,8 @@ curl -s http://:9003/api/v1/tenant_quotas # Query one tenant curl -s "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" -# Upsert an explicit policy. Explicit tenant policies must be positive. +# Upsert an explicit policy. Explicit tenant policies must be between 1 byte +# and 2^63 - 1 bytes. curl -s -X PUT "http://:9003/api/v1/tenant_quotas?tenant_id=tenant-a" \ -H 'Content-Type: application/json' \ -d '{"requested_quota_bytes":2147483648}' @@ -73,16 +74,16 @@ Each tenant quota snapshot returns: "tenant_id": "tenant-a", "requested_quota_bytes": 2147483648, "effective_quota_bytes": 2147483648, - "used_bytes": 0, - "reserved_bytes": 0, - "committed_count": 0, - "metadata_object_count": 0, + "charged_bytes": 0, + "admission_closed": false, "over_quota": false, "has_explicit_policy": true } } ``` +`charged_bytes` includes completed MEMORY replicas and in-flight MEMORY allocations. Put, Copy, Move, and promotion charge quota when admission starts; failed, revoked, partially completed, or expired operations refund the unused charge. `admission_closed` is `true` when the account rejects new writes, including after its explicit policy is removed. + In HA mode, quota admin requests are served only by the active master service. Standby, candidate, or inactive services return HTTP 503. If strict multi-tenant mode is disabled, the quota admin API returns HTTP 409 with `UNAVAILABLE_IN_CURRENT_MODE`. Deleting a non-empty tenant returns HTTP 409 with `TENANT_NOT_EMPTY`. ## SGLang diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index 4e4270d75d..a58048f83d 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -32,6 +32,7 @@ #include "master_metric_manager.h" #include "mutex.h" #include "segment.h" +#include "tenant_quota_ledger.h" #include "tenant_quota_sharded.h" #include "tenant_quota_policy_store.h" #include "types.h" @@ -1032,9 +1033,7 @@ class MasterService { const bool hard_pinned{false}; // immutable, set at creation bool memory_cache_total_accounted{false}; bool disk_cache_total_accounted{false}; - uint64_t reserved_quota_charge_bytes{0}; - uint64_t committed_quota_charge_bytes{0}; - uint64_t pending_replaced_quota_charge_bytes{0}; + TenantQuotaLedger quota_ledger; void AddReplicas(std::vector&& replicas) { replicas_.insert(replicas_.end(), @@ -1413,7 +1412,7 @@ class MasterService { } type; ReplicaID source_id; std::vector replica_ids; - uint64_t reserved_quota_charge_bytes{0}; + uint64_t pending_quota_charge_bytes{0}; }; struct OffloadingTask { @@ -1477,7 +1476,7 @@ class MasterService { ReplicaID source_id; // the LOCAL_DISK replica being promoted ReplicaID alloc_id{0}; // the new MEMORY replica staged by AllocStart uint64_t object_size; - uint64_t reserved_quota_charge_bytes{0}; + uint64_t pending_quota_charge_bytes{0}; std::chrono::system_clock::time_point start_time; UUID holder_id; // owner of source LOCAL_DISK; only Notifier allowed }; @@ -1485,6 +1484,7 @@ class MasterService { static constexpr size_t kNumShards = 1024; // Number of metadata shards struct TenantState { + TenantQuotaHandle quota_account{nullptr}; std::unordered_map metadata; std::unordered_set processing_keys; std::unordered_map @@ -1717,21 +1717,19 @@ class MasterService { TenantState& tenant_state, std::unordered_map::iterator it, const TenantId& tenant_id, QuotaEraseMode quota_mode); + tl::expected SettlePrimaryWriteQuotaIfReady( + TenantState& tenant_state, ObjectMetadata& metadata); uint64_t CompletedMemoryQuotaCharge(const ObjectMetadata& metadata) const; uint64_t RequestedMemoryQuotaCharge(uint64_t value_length, const ReplicateConfig& config) const; - bool ShouldProtectZeroChargeMetadataCreate( - uint64_t requested_quota_charge) const; - tl::expected ReserveTenantQuota(const TenantId& tenant_id, - uint64_t bytes); - void CommitTenantQuota(const TenantId& tenant_id, uint64_t bytes); - void AbortTenantQuota(const TenantId& tenant_id, uint64_t bytes); - void ReleaseTenantQuota(const TenantId& tenant_id, uint64_t bytes); - void ReleaseTenantQuotaPartial(const TenantId& tenant_id, uint64_t bytes); - void CommitAdditionalTenantQuota(const TenantId& tenant_id, uint64_t bytes); - void IncrementTenantMetadataObjectCount(const TenantId& tenant_id); - void DecrementTenantMetadataObjectCount(const TenantId& tenant_id); - void ReleaseCommittedQuotaCharge(ObjectMetadata& metadata, uint64_t bytes); + TenantState& GetOrCreateTenantState(MetadataShard& shard, + const TenantId& tenant_id); + TenantQuotaHandle GetBoundTenantQuotaHandle( + const TenantState& tenant_state) const; + tl::expected ChargeTenantQuota( + TenantQuotaHandle account, uint64_t bytes, + uint64_t* deficit_bytes = nullptr); + void ReleaseTenantQuota(TenantQuotaHandle account, uint64_t bytes); void RecomputeTenantEffectiveQuotas(); void RebuildTenantQuotaUsageFromMetadata(); void LoadTenantQuotaPoliciesFromStoreOrThrow(); @@ -1799,6 +1797,7 @@ class MasterService { const TenantId& tenant_id, const std::chrono::system_clock::time_point& now, const ResolvedSoftPinRequest& soft_pin_request, + uint64_t& quota_deficit_bytes, std::optional committed_soft_pin_timeout = std::nullopt) -> tl::expected, ErrorCode>; @@ -1874,16 +1873,17 @@ class MasterService { size_t CountCandidatesForTesting(const TenantId& tenant_id); void ResetCandidateBackoffsForTesting(); - // Erase any in-flight PromotionTask for `key`, abort any staged promotion - // quota reservation, and decrement the cluster-wide in-flight counter. Safe - // no-op if no task exists. - void ErasePromotionTaskIfPresent( - TenantState& tenant_state, const std::string& key, - const TenantId& tenant_id) NO_THREAD_SAFETY_ANALYSIS { + // Erase any in-flight PromotionTask for `key`, refund its pending charge, + // and decrement the cluster-wide in-flight counter. Safe no-op if no task + // exists. + void ErasePromotionTaskIfPresent(TenantState& tenant_state, + const std::string& key) + NO_THREAD_SAFETY_ANALYSIS { auto task_it = tenant_state.promotion_tasks.find(key); if (task_it != tenant_state.promotion_tasks.end()) { - AbortTenantQuota(tenant_id, - task_it->second.reserved_quota_charge_bytes); + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), + std::exchange(task_it->second.pending_quota_charge_bytes, 0)); tenant_state.promotion_tasks.erase(task_it); promotion_in_flight_.fetch_sub(1, std::memory_order_relaxed); MasterMetricManager::instance().dec_promotion_in_flight(); @@ -1954,6 +1954,9 @@ class MasterService { ? ReplicationTaskIterator{} : tenant_state_->replication_tasks.find( object_id_.user_key)) { + if (tenant_state_ != nullptr) { + service_->GetBoundTenantQuotaHandle(*tenant_state_); + } // Automatically clean up invalid handles (memory replicas only). // Note: We only check memory replicas here to avoid lock order // violation (client_mutex_ must be acquired before metadata shard). @@ -1978,9 +1981,19 @@ class MasterService { *tenant_state_, it_->second, removed_replica_ids); const uint64_t after_charge = service_->CompletedMemoryQuotaCharge(it_->second); - if (before_charge > after_charge) { - service_->ReleaseCommittedQuotaCharge( - it_->second, before_charge - after_charge); + if (service_->enable_multi_tenants_ && + before_charge > after_charge) { + auto release_result = + it_->second.quota_ledger.ReleaseCommitted( + service_->GetBoundTenantQuotaHandle(*tenant_state_), + before_charge - after_charge); + if (!release_result) { + LOG(ERROR) + << "tenant quota committed release mismatch tenant=" + << object_id_.tenant_id.value() + << ", key=" << object_id_.user_key + << ", bytes=" << before_charge - after_charge; + } } // If no valid replicas remain, delete the whole object. if (!it_->second.IsValid()) { @@ -1992,8 +2005,7 @@ class MasterService { this->Erase(); if (tenant_state_ != nullptr) { service_->ErasePromotionTaskIfPresent( - *tenant_state_, object_id_.user_key, - object_id_.tenant_id); + *tenant_state_, object_id_.user_key); MaybeEraseEmptyTenant(); } } @@ -2071,10 +2083,6 @@ class MasterService { std::nullopt, enable_hard_pin, data_type, group_id, object_id_.tenant_id, object_id_.user_key)); it_ = result.first; - if (result.second) { - service_->IncrementTenantMetadataObjectCount( - object_id_.tenant_id); - } } private: @@ -2088,10 +2096,9 @@ class MasterService { if (tenant_state_ != nullptr) { return; } - auto result = - shard_guard_->tenants.try_emplace(object_id_.tenant_id); - tenant_it_ = result.first; - tenant_state_ = &tenant_it_->second; + tenant_state_ = &service_->GetOrCreateTenantState( + shard_guard_.get(), object_id_.tenant_id); + tenant_it_ = shard_guard_->tenants.find(object_id_.tenant_id); it_ = tenant_state_->metadata.end(); processing_it_ = tenant_state_->processing_keys.end(); replication_task_it_ = tenant_state_->replication_tasks.end(); diff --git a/mooncake-store/include/tenant_quota.h b/mooncake-store/include/tenant_quota.h index 6621e54974..3e89542c75 100644 --- a/mooncake-store/include/tenant_quota.h +++ b/mooncake-store/include/tenant_quota.h @@ -1,8 +1,10 @@ #pragma once +#include #include #include #include +#include #include #include #include @@ -18,23 +20,15 @@ struct TenantQuotaSnapshot { TenantId tenant_id; uint64_t requested_quota_bytes = 0; uint64_t effective_quota_bytes = 0; - uint64_t used_bytes = 0; - uint64_t reserved_bytes = 0; - uint64_t committed_count = 0; - uint64_t metadata_object_count = 0; + uint64_t charged_bytes = 0; + bool admission_closed = true; bool has_explicit_policy = false; bool over_quota = false; }; -struct TenantQuotaUsage { - uint64_t used_bytes = 0; - uint64_t committed_count = 0; - uint64_t metadata_object_count = 0; -}; - using TenantQuotaPolicyMap = std::map; using TenantQuotaUsageMap = - std::unordered_map; + std::unordered_map; enum class TenantQuotaError { kQuotaExceeded, @@ -45,74 +39,93 @@ enum class TenantQuotaError { kTenantNotFound, }; +struct TenantQuotaChargeFailure { + TenantQuotaError error; + uint64_t deficit_bytes = 0; +}; + using TenantQuotaResult = tl::expected; -using TenantQuotaPolicyResult = tl::expected; +using TenantQuotaChargeResult = tl::expected; + +class TenantQuotaAccount { + public: + static constexpr uint64_t kAdmissionClosed = 1ULL << 63; + static constexpr uint64_t kChargedBytesMask = kAdmissionClosed - 1; + static constexpr uint64_t kMaxChargedBytes = kChargedBytesMask; + + TenantQuotaChargeResult TryCharge(uint64_t bytes); + TenantQuotaResult Release(uint64_t bytes); + + uint64_t ChargedBytes() const; + uint64_t EffectiveQuotaBytes() const; + bool AdmissionClosed() const; + + private: + friend class TenantQuotaTable; + + void BeginPolicyUpdate(); + void EndPolicyUpdate(); + void SetAdmissionClosed(bool closed); + void ApplyEffectiveQuota(uint64_t effective_quota_bytes); + void SetChargedBytesForRebuild(uint64_t charged_bytes); + + // Bit 63 closes admission; bits 0..62 contain charged bytes. + alignas(64) std::atomic charged_state_{kAdmissionClosed}; + + // Keep control-plane writes off the charged-state cache line. + alignas(64) std::atomic effective_quota_bytes_{0}; + std::atomic policy_sequence_{0}; + + // Accessed only under external control-plane synchronization. + uint64_t requested_quota_bytes_{0}; + bool has_explicit_policy_{false}; +}; + +using TenantQuotaHandle = TenantQuotaAccount*; template class ShardedTenantQuotaTable; -// Single-threaded tenant quota state machine. This class owns quota rules and -// accounting invariants, but deliberately contains no locking or sharding. +// Control-plane registry for stable quota accounts. Callers must provide +// external synchronization. Charge and release bypass this table and operate +// directly on TenantQuotaHandle. class TenantQuotaTable { public: TenantQuotaResult UpsertTenantPolicy(const TenantId& tenant_id, uint64_t requested_quota_bytes); - TenantQuotaPolicyResult DisableTenantPolicyIfEmpty( - const TenantId& tenant_id); - void ApplyTenantPolicies(const TenantQuotaPolicyMap& policies); + TenantQuotaResult DisableTenantPolicyIfEmpty(const TenantId& tenant_id); + TenantQuotaResult ApplyTenantPolicies(const TenantQuotaPolicyMap& policies); TenantQuotaPolicyMap GetTenantPolicies() const; void RecomputeEffectiveQuotas(uint64_t allocatable_capacity_bytes); bool IsTenantRegistered(const TenantId& tenant_id) const; + // May create a stable closed tombstone for a previously unseen tenant. + TenantQuotaHandle GetOrCreateTenantHandle(const TenantId& tenant_id); std::optional GetTenantSnapshot( const TenantId& tenant_id) const; std::vector ListTenantSnapshots() const; - uint64_t ComputeDeficit(const TenantId& tenant_id, - uint64_t incoming_bytes) const; - TenantQuotaResult Reserve(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult Commit(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult CommitAdditional(const TenantId& tenant_id, - uint64_t bytes); - TenantQuotaResult Abort(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult Release(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult ReleasePartial(const TenantId& tenant_id, uint64_t bytes); - - void IncrementMetadataObjectCount(const TenantId& tenant_id); - TenantQuotaResult DecrementMetadataObjectCount(const TenantId& tenant_id); - void RebuildUsage(const TenantQuotaUsageMap& usage); + // Rebuild overwrites runtime accounting and must only run while data-plane + // charge/release operations are quiescent. + TenantQuotaResult RebuildUsage(const TenantQuotaUsageMap& usage); private: template friend class ShardedTenantQuotaTable; - struct TenantQuotaState { - uint64_t requested_quota_bytes = 0; - uint64_t effective_quota_bytes = 0; - uint64_t used_bytes = 0; - uint64_t reserved_bytes = 0; - uint64_t committed_count = 0; - uint64_t metadata_object_count = 0; - bool has_explicit_policy = false; - bool over_quota = false; - }; - - using StateMap = std::map; + using AccountMap = std::map>; - TenantQuotaState& GetOrCreateState(const TenantId& tenant_id); + TenantQuotaAccount& GetOrCreateAccount(const TenantId& tenant_id); TenantQuotaSnapshot MakeSnapshot(const TenantId& tenant_id, - const TenantQuotaState& state) const; - static bool IsLazyEmptyTenant(const TenantQuotaState& state); - static void RefreshOverQuota(TenantQuotaState* state); + const TenantQuotaAccount& account) const; static std::map BuildEffectiveQuotaAssignments( const std::vector& tenants, uint64_t allocatable_capacity_bytes); void ApplyEffectiveQuotas( const std::map& effective_quotas); - void EraseIfLazyEmpty(StateMap::iterator it); - StateMap tenants_; + AccountMap accounts_; }; } // namespace mooncake diff --git a/mooncake-store/include/tenant_quota_ledger.h b/mooncake-store/include/tenant_quota_ledger.h new file mode 100644 index 0000000000..78416e8dac --- /dev/null +++ b/mooncake-store/include/tenant_quota_ledger.h @@ -0,0 +1,65 @@ +#pragma once + +#include + +#include "tenant_quota.h" + +namespace mooncake { + +// Per-object accounting state for quota bytes already charged to a stable +// TenantQuotaAccount. The ledger never owns the account and never releases +// quota implicitly; every mutation receives the account handle explicitly. +class TenantQuotaLedger { + public: + TenantQuotaLedger() = default; + TenantQuotaLedger(const TenantQuotaLedger&) = delete; + TenantQuotaLedger& operator=(const TenantQuotaLedger&) = delete; + TenantQuotaLedger(TenantQuotaLedger&&) = delete; + TenantQuotaLedger& operator=(TenantQuotaLedger&&) = delete; + ~TenantQuotaLedger() = default; + + // Records bytes that the caller has already charged with TryCharge(). + TenantQuotaResult AdoptPendingCharge(TenantQuotaHandle account, + uint64_t bytes); + + // Completes a primary Put/Upsert lifecycle. The replacement bucket is + // always released; pending/committed bytes are reconciled to + // actual_committed_bytes. Validation and the account release complete + // before any local field is changed. + TenantQuotaResult SettlePrimaryWrite(TenantQuotaHandle account, + uint64_t actual_committed_bytes); + + // Settles a temporary task charge into this object's committed bytes. The + // caller erases or clears the task field only after this succeeds. + TenantQuotaResult SettleAdditional(TenantQuotaHandle account, + uint64_t pending_bytes, + uint64_t actual_bytes); + + TenantQuotaResult RefundPending(TenantQuotaHandle account); + TenantQuotaResult ReleaseCommitted(TenantQuotaHandle account, + uint64_t bytes); + + // Transfers this ledger's committed/replacement contribution into the + // destination's replacement bucket without changing global charged bytes. + TenantQuotaResult TransferReplacementCharge(TenantQuotaHandle account, + TenantQuotaLedger& destination); + TenantQuotaResult ReleaseReplacement(TenantQuotaHandle account); + TenantQuotaResult ReleaseAll(TenantQuotaHandle account); + + // Rebuild is only valid while data-plane charge/release is quiescent. It + // changes local state only; the caller rebuilds the account separately. + TenantQuotaResult Rebuild(TenantQuotaHandle account, + uint64_t committed_bytes); + + uint64_t PendingBytes() const { return pending_bytes_; } + uint64_t CommittedBytes() const { return committed_bytes_; } + uint64_t ReplacedBytes() const { return replaced_bytes_; } + uint64_t TotalChargedBytes() const; + + private: + uint64_t pending_bytes_{0}; + uint64_t committed_bytes_{0}; + uint64_t replaced_bytes_{0}; +}; + +} // namespace mooncake diff --git a/mooncake-store/include/tenant_quota_sharded.h b/mooncake-store/include/tenant_quota_sharded.h index afe0e47e06..29f294c1ce 100644 --- a/mooncake-store/include/tenant_quota_sharded.h +++ b/mooncake-store/include/tenant_quota_sharded.h @@ -20,33 +20,24 @@ class ShardedTenantQuotaTable { TenantQuotaResult UpsertTenantPolicy(const TenantId& tenant_id, uint64_t requested_quota_bytes, uint64_t allocatable_capacity_bytes); - TenantQuotaPolicyResult DisableTenantPolicyIfEmpty( - const TenantId& tenant_id); - void ApplyTenantPolicies(const TenantQuotaPolicyMap& policies, - uint64_t allocatable_capacity_bytes); + TenantQuotaResult DisableTenantPolicyIfEmpty(const TenantId& tenant_id); + TenantQuotaResult ApplyTenantPolicies(const TenantQuotaPolicyMap& policies, + uint64_t allocatable_capacity_bytes); TenantQuotaPolicyMap GetTenantPolicies() const; void RecomputeEffectiveQuotas(uint64_t allocatable_capacity_bytes); bool IsTenantRegistered(const TenantId& tenant_id) const; + // May create a stable closed tombstone for a previously unseen tenant. + TenantQuotaHandle GetOrCreateTenantHandle(const TenantId& tenant_id); std::optional GetTenantSnapshot( const TenantId& tenant_id) const; std::vector ListTenantSnapshots() const; - uint64_t ComputeDeficit(const TenantId& tenant_id, - uint64_t incoming_bytes) const; - TenantQuotaResult Reserve(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult Commit(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult CommitAdditional(const TenantId& tenant_id, - uint64_t bytes); - TenantQuotaResult Abort(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult Release(const TenantId& tenant_id, uint64_t bytes); - TenantQuotaResult ReleasePartial(const TenantId& tenant_id, uint64_t bytes); - - void IncrementMetadataObjectCount(const TenantId& tenant_id); - TenantQuotaResult DecrementMetadataObjectCount(const TenantId& tenant_id); - void RebuildUsage(const TenantQuotaUsageMap& usage, - uint64_t allocatable_capacity_bytes); + // Rebuild overwrites runtime accounting and must only run while data-plane + // charge/release operations are quiescent. + TenantQuotaResult RebuildUsage(const TenantQuotaUsageMap& usage, + uint64_t allocatable_capacity_bytes); private: struct Shard { diff --git a/mooncake-store/include/tenant_quota_sharded_impl.h b/mooncake-store/include/tenant_quota_sharded_impl.h index d74c272a4d..2b938d3ea3 100644 --- a/mooncake-store/include/tenant_quota_sharded_impl.h +++ b/mooncake-store/include/tenant_quota_sharded_impl.h @@ -25,7 +25,7 @@ TenantQuotaResult ShardedTenantQuotaTable::UpsertTenantPolicy( } template -TenantQuotaPolicyResult +TenantQuotaResult ShardedTenantQuotaTable::DisableTenantPolicyIfEmpty( const TenantId& tenant_id) { std::lock_guard recompute_lock(recompute_mutex_); @@ -35,8 +35,15 @@ ShardedTenantQuotaTable::DisableTenantPolicyIfEmpty( } template -void ShardedTenantQuotaTable::ApplyTenantPolicies( +TenantQuotaResult ShardedTenantQuotaTable::ApplyTenantPolicies( const TenantQuotaPolicyMap& policies, uint64_t allocatable_capacity_bytes) { + for (const auto& [_, requested_quota_bytes] : policies) { + if (requested_quota_bytes == 0 || + requested_quota_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return tl::make_unexpected(TenantQuotaError::kInvalidArgument); + } + } + std::lock_guard recompute_lock(recompute_mutex_); std::array grouped_policies; for (const auto& [tenant_id, requested_quota_bytes] : policies) { @@ -47,9 +54,13 @@ void ShardedTenantQuotaTable::ApplyTenantPolicies( for (size_t i = 0; i < kNumShards; ++i) { auto& shard = shards_[i]; std::lock_guard lock(shard.mutex); - shard.table.ApplyTenantPolicies(grouped_policies[i]); + auto result = shard.table.ApplyTenantPolicies(grouped_policies[i]); + if (!result) { + return result; + } } RecomputeEffectiveQuotasLocked(allocatable_capacity_bytes); + return {}; } template @@ -79,6 +90,14 @@ bool ShardedTenantQuotaTable::IsTenantRegistered( return shard.table.IsTenantRegistered(tenant_id); } +template +TenantQuotaHandle ShardedTenantQuotaTable::GetOrCreateTenantHandle( + const TenantId& tenant_id) { + auto& shard = GetShard(tenant_id); + std::lock_guard lock(shard.mutex); + return shard.table.GetOrCreateTenantHandle(tenant_id); +} + template std::optional ShardedTenantQuotaTable::GetTenantSnapshot( @@ -106,94 +125,31 @@ ShardedTenantQuotaTable::ListTenantSnapshots() const { } template -uint64_t ShardedTenantQuotaTable::ComputeDeficit( - const TenantId& tenant_id, uint64_t incoming_bytes) const { - const auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.ComputeDeficit(tenant_id, incoming_bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::Reserve( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.Reserve(tenant_id, bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::Commit( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.Commit(tenant_id, bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::CommitAdditional( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.CommitAdditional(tenant_id, bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::Abort( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.Abort(tenant_id, bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::Release( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.Release(tenant_id, bytes); -} - -template -TenantQuotaResult ShardedTenantQuotaTable::ReleasePartial( - const TenantId& tenant_id, uint64_t bytes) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.ReleasePartial(tenant_id, bytes); -} - -template -void ShardedTenantQuotaTable::IncrementMetadataObjectCount( - const TenantId& tenant_id) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - shard.table.IncrementMetadataObjectCount(tenant_id); -} - -template -TenantQuotaResult -ShardedTenantQuotaTable::DecrementMetadataObjectCount( - const TenantId& tenant_id) { - auto& shard = GetShard(tenant_id); - std::lock_guard lock(shard.mutex); - return shard.table.DecrementMetadataObjectCount(tenant_id); -} - -template -void ShardedTenantQuotaTable::RebuildUsage( +TenantQuotaResult ShardedTenantQuotaTable::RebuildUsage( const TenantQuotaUsageMap& usage, uint64_t allocatable_capacity_bytes) { + for (const auto& [_, charged_bytes] : usage) { + if (charged_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return tl::make_unexpected(TenantQuotaError::kInvalidArgument); + } + } + std::lock_guard recompute_lock(recompute_mutex_); std::array grouped_usage; - for (const auto& [tenant_id, tenant_usage] : usage) { + for (const auto& [tenant_id, charged_bytes] : usage) { grouped_usage[GetShardIndex(tenant_id)].emplace(tenant_id, - tenant_usage); + charged_bytes); } for (size_t i = 0; i < kNumShards; ++i) { auto& shard = shards_[i]; std::lock_guard lock(shard.mutex); - shard.table.RebuildUsage(grouped_usage[i]); + auto result = shard.table.RebuildUsage(grouped_usage[i]); + if (!result) { + return result; + } } RecomputeEffectiveQuotasLocked(allocatable_capacity_bytes); + return {}; } template diff --git a/mooncake-store/src/CMakeLists.txt b/mooncake-store/src/CMakeLists.txt index 6fbf2f1a9a..b747a37cb9 100644 --- a/mooncake-store/src/CMakeLists.txt +++ b/mooncake-store/src/CMakeLists.txt @@ -19,6 +19,7 @@ set(MOONCAKE_STORE_SOURCES segment.cpp transfer_task.cpp tenant_quota.cpp + tenant_quota_ledger.cpp tenant_quota_policy_store.cpp rpc_service.cpp master_admin_service.cpp diff --git a/mooncake-store/src/master_admin_service.cpp b/mooncake-store/src/master_admin_service.cpp index 2b4d3fbf77..46d61d32d7 100644 --- a/mooncake-store/src/master_admin_service.cpp +++ b/mooncake-store/src/master_admin_service.cpp @@ -175,16 +175,14 @@ struct HttpTenantQuotaSnapshot { std::string tenant_id; uint64_t requested_quota_bytes{0}; uint64_t effective_quota_bytes{0}; - uint64_t used_bytes{0}; - uint64_t reserved_bytes{0}; - uint64_t committed_count{0}; - uint64_t metadata_object_count{0}; + uint64_t charged_bytes{0}; + bool admission_closed{true}; bool over_quota{false}; bool has_explicit_policy{false}; }; YLT_REFL(HttpTenantQuotaSnapshot, tenant_id, requested_quota_bytes, - effective_quota_bytes, used_bytes, reserved_bytes, committed_count, - metadata_object_count, over_quota, has_explicit_policy); + effective_quota_bytes, charged_bytes, admission_closed, over_quota, + has_explicit_policy); HttpTenantQuotaSnapshot ToHttpTenantQuotaSnapshot( const TenantQuotaSnapshot& snapshot) { @@ -192,10 +190,8 @@ HttpTenantQuotaSnapshot ToHttpTenantQuotaSnapshot( .tenant_id = snapshot.tenant_id.value(), .requested_quota_bytes = snapshot.requested_quota_bytes, .effective_quota_bytes = snapshot.effective_quota_bytes, - .used_bytes = snapshot.used_bytes, - .reserved_bytes = snapshot.reserved_bytes, - .committed_count = snapshot.committed_count, - .metadata_object_count = snapshot.metadata_object_count, + .charged_bytes = snapshot.charged_bytes, + .admission_closed = snapshot.admission_closed, .over_quota = snapshot.over_quota, .has_explicit_policy = snapshot.has_explicit_policy, }; @@ -393,18 +389,12 @@ std::string MasterAdminServer::BuildTenantQuotaMetricsText() const { << "# HELP mooncake_tenant_quota_effective_bytes Effective tenant " "quota in bytes\n" << "# TYPE mooncake_tenant_quota_effective_bytes gauge\n" - << "# HELP mooncake_tenant_quota_used_bytes Tenant committed quota " - "usage in bytes\n" - << "# TYPE mooncake_tenant_quota_used_bytes gauge\n" - << "# HELP mooncake_tenant_quota_reserved_bytes Tenant reserved quota " - "usage in bytes\n" - << "# TYPE mooncake_tenant_quota_reserved_bytes gauge\n" - << "# HELP mooncake_tenant_quota_committed_count Tenant committed " - "object count\n" - << "# TYPE mooncake_tenant_quota_committed_count gauge\n" - << "# HELP mooncake_tenant_quota_metadata_object_count Tenant " - "metadata object count\n" - << "# TYPE mooncake_tenant_quota_metadata_object_count gauge\n" + << "# HELP mooncake_tenant_quota_charged_bytes Tenant quota charge in " + "bytes\n" + << "# TYPE mooncake_tenant_quota_charged_bytes gauge\n" + << "# HELP mooncake_tenant_quota_admission_closed Tenant quota " + "admission-closed flag\n" + << "# TYPE mooncake_tenant_quota_admission_closed gauge\n" << "# HELP mooncake_tenant_quota_over_quota Tenant over-quota flag\n" << "# TYPE mooncake_tenant_quota_over_quota gauge\n" << "# HELP mooncake_tenant_quota_explicit_policy Tenant explicit " @@ -420,15 +410,11 @@ std::string MasterAdminServer::BuildTenantQuotaMetricsText() const { tenant_metrics << "mooncake_tenant_quota_effective_bytes{tenant_id=\"" << tenant << "\"} " << snapshot.effective_quota_bytes << "\n"; - tenant_metrics << "mooncake_tenant_quota_used_bytes{tenant_id=\"" - << tenant << "\"} " << snapshot.used_bytes << "\n"; - tenant_metrics << "mooncake_tenant_quota_reserved_bytes{tenant_id=\"" - << tenant << "\"} " << snapshot.reserved_bytes << "\n"; - tenant_metrics << "mooncake_tenant_quota_committed_count{tenant_id=\"" - << tenant << "\"} " << snapshot.committed_count << "\n"; - tenant_metrics - << "mooncake_tenant_quota_metadata_object_count{tenant_id=\"" - << tenant << "\"} " << snapshot.metadata_object_count << "\n"; + tenant_metrics << "mooncake_tenant_quota_charged_bytes{tenant_id=\"" + << tenant << "\"} " << snapshot.charged_bytes << "\n"; + tenant_metrics << "mooncake_tenant_quota_admission_closed{tenant_id=\"" + << tenant << "\"} " + << (snapshot.admission_closed ? 1 : 0) << "\n"; tenant_metrics << "mooncake_tenant_quota_over_quota{tenant_id=\"" << tenant << "\"} " << (snapshot.over_quota ? 1 : 0) << "\n"; @@ -1111,10 +1097,12 @@ void MasterAdminServer::HandleUpsertTenantQuota( ErrorCode::INVALID_PARAMS, body_result.error()); return; } - if (body_result->requested_quota_bytes == 0) { + if (body_result->requested_quota_bytes == 0 || + body_result->requested_quota_bytes > + TenantQuotaAccount::kMaxChargedBytes) { WriteErrorResponse(resp, coro_http::status_type::bad_request, ErrorCode::INVALID_PARAMS, - "Tenant quota must be positive"); + "Tenant quota must be in [1, 2^63 - 1] bytes"); return; } diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index c61a3af482..2bd9d0bc0f 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -133,6 +133,18 @@ bool HasExpectedReplicaAllocation(const ReplicateConfig& config, allocated_nof_replicas == config.nof_replica_num; } +void LogTenantQuotaLedgerError(const TenantQuotaResult& result, + std::string_view operation, + const TenantId& tenant_id, + std::string_view key) { + if (result) { + return; + } + LOG(ERROR) << "tenant quota ledger error operation=" << operation + << ", tenant=" << tenant_id.value() << ", key=" << key + << ", error=" << static_cast(result.error()); +} + tl::expected GetGroupIdForKey( const ReplicateConfig& config, size_t key_count, size_t key_index) { if (!config.group_ids.has_value()) { @@ -690,7 +702,8 @@ MasterService::UpsertTenantQuotaPolicy(const TenantId& tenant_id, if (!enable_multi_tenants_) { return tl::make_unexpected(ErrorCode::UNAVAILABLE_IN_CURRENT_MODE); } - if (requested_quota_bytes == 0) { + if (requested_quota_bytes == 0 || + requested_quota_bytes > TenantQuotaAccount::kMaxChargedBytes) { return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } @@ -747,13 +760,7 @@ MasterService::DeleteTenantQuotaPolicy(const TenantId& tenant_id) { : ErrorCode::OBJECT_NOT_FOUND); } - auto post_mark_snapshot = GetTenantQuotaSnapshot(tenant_id); - if (TenantHasObjects(tenant_id) || - (post_mark_snapshot.has_value() && - (post_mark_snapshot->used_bytes != 0 || - post_mark_snapshot->reserved_bytes != 0 || - post_mark_snapshot->committed_count != 0 || - post_mark_snapshot->metadata_object_count != 0))) { + if (TenantHasObjects(tenant_id)) { restore_policy(); return tl::make_unexpected(ErrorCode::TENANT_NOT_EMPTY); } @@ -1300,7 +1307,11 @@ void MasterService::ApplyTenantQuotaPolicies( } std::lock_guard recompute_lock(tenant_quota_recompute_mutex_); const uint64_t capacity = GetTenantQuotaAllocatableCapacityBytes(); - tenant_quota_table_.ApplyTenantPolicies(policies, capacity); + auto result = tenant_quota_table_.ApplyTenantPolicies(policies, capacity); + if (!result) { + throw std::invalid_argument( + "tenant quota policy exceeds atomic accounting range"); + } } void MasterService::LoadTenantQuotaPoliciesFromStoreOrThrow() { @@ -1322,10 +1333,15 @@ void MasterService::LoadTenantQuotaPoliciesFromStoreOrThrow() { uint64_t MasterService::CompletedMemoryQuotaCharge( const ObjectMetadata& metadata) const { - return static_cast(metadata.size) * - metadata.CountReplicas([](const Replica& replica) { - return replica.is_memory_replica() && replica.is_completed(); - }); + const auto completed_replicas = + metadata.CountReplicas([](const Replica& replica) { + return replica.is_memory_replica() && replica.is_completed(); + }); + const unsigned __int128 charge = + static_cast(metadata.size) * completed_replicas; + return charge > std::numeric_limits::max() + ? std::numeric_limits::max() + : static_cast(charge); } uint64_t MasterService::RequestedMemoryQuotaCharge( @@ -1338,11 +1354,6 @@ uint64_t MasterService::RequestedMemoryQuotaCharge( return static_cast(charge); } -bool MasterService::ShouldProtectZeroChargeMetadataCreate( - uint64_t requested_quota_charge) const { - return enable_multi_tenants_ && requested_quota_charge == 0; -} - uint64_t MasterService::GetTenantQuotaAllocatableCapacityBytes() { uint64_t capacity = 0; ScopedSegmentAccess segment_access = segment_manager_.getSegmentAccess(); @@ -1368,110 +1379,67 @@ void MasterService::RecomputeTenantEffectiveQuotas() { tenant_quota_table_.RecomputeEffectiveQuotas(capacity); } -tl::expected MasterService::ReserveTenantQuota( - const TenantId& tenant_id, uint64_t bytes) { +MasterService::TenantState& MasterService::GetOrCreateTenantState( + MetadataShard& shard, const TenantId& tenant_id) { + auto it = shard.tenants.try_emplace(tenant_id).first; + if (enable_multi_tenants_ && it->second.quota_account == nullptr) { + it->second.quota_account = + tenant_quota_table_.GetOrCreateTenantHandle(tenant_id); + } + return it->second; +} + +TenantQuotaHandle MasterService::GetBoundTenantQuotaHandle( + const TenantState& tenant_state) const { + if (!enable_multi_tenants_) { + return nullptr; + } + assert(tenant_state.quota_account != nullptr); + return tenant_state.quota_account; +} + +tl::expected MasterService::ChargeTenantQuota( + TenantQuotaHandle account, uint64_t bytes, uint64_t* deficit_bytes) { if (!enable_multi_tenants_) { return {}; } - auto result = tenant_quota_table_.Reserve(tenant_id, bytes); + if (account == nullptr) { + LOG(ERROR) << "tenant quota charge attempted without a bound handle"; + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + auto result = account->TryCharge(bytes); if (result) { + if (deficit_bytes != nullptr) { + *deficit_bytes = 0; + } return {}; } + if (deficit_bytes != nullptr) { + *deficit_bytes = result.error().deficit_bytes; + } return tl::make_unexpected( - result.error() == TenantQuotaError::kTenantNotRegistered + result.error().error == TenantQuotaError::kTenantNotRegistered ? ErrorCode::TENANT_NOT_REGISTERED - : result.error() == TenantQuotaError::kQuotaExceeded + : result.error().error == TenantQuotaError::kQuotaExceeded ? ErrorCode::TENANT_QUOTA_EXCEEDED + : result.error().error == TenantQuotaError::kInvalidArgument + ? ErrorCode::INVALID_PARAMS : ErrorCode::INTERNAL_ERROR); } -void MasterService::CommitTenantQuota(const TenantId& tenant_id, - uint64_t bytes) { - if (!enable_multi_tenants_ || bytes == 0) { - return; - } - if (!tenant_quota_table_.Commit(tenant_id, bytes)) { - LOG(ERROR) << "tenant quota commit mismatch tenant=" - << tenant_id.value() << ", bytes=" << bytes; - } -} - -void MasterService::AbortTenantQuota(const TenantId& tenant_id, - uint64_t bytes) { - if (!enable_multi_tenants_ || bytes == 0) { - return; - } - if (!tenant_quota_table_.Abort(tenant_id, bytes)) { - LOG(ERROR) << "tenant quota abort mismatch tenant=" << tenant_id.value() - << ", bytes=" << bytes; - } -} - -void MasterService::ReleaseTenantQuota(const TenantId& tenant_id, +void MasterService::ReleaseTenantQuota(TenantQuotaHandle account, uint64_t bytes) { if (!enable_multi_tenants_ || bytes == 0) { return; } - if (!tenant_quota_table_.Release(tenant_id, bytes)) { - LOG(ERROR) << "tenant quota release mismatch tenant=" - << tenant_id.value() << ", bytes=" << bytes; - } -} - -void MasterService::ReleaseTenantQuotaPartial(const TenantId& tenant_id, - uint64_t bytes) { - if (!enable_multi_tenants_ || bytes == 0) { - return; - } - if (!tenant_quota_table_.ReleasePartial(tenant_id, bytes)) { - LOG(ERROR) << "tenant quota partial release mismatch tenant=" - << tenant_id.value() << ", bytes=" << bytes; - } -} - -void MasterService::CommitAdditionalTenantQuota(const TenantId& tenant_id, - uint64_t bytes) { - if (!enable_multi_tenants_ || bytes == 0) { - return; - } - if (!tenant_quota_table_.CommitAdditional(tenant_id, bytes)) { - LOG(ERROR) << "tenant quota additional commit mismatch tenant=" - << tenant_id.value() << ", bytes=" << bytes; - } -} - -void MasterService::IncrementTenantMetadataObjectCount( - const TenantId& tenant_id) { - if (!enable_multi_tenants_) { - return; - } - tenant_quota_table_.IncrementMetadataObjectCount(tenant_id); -} - -void MasterService::DecrementTenantMetadataObjectCount( - const TenantId& tenant_id) { - if (!enable_multi_tenants_) { - return; - } - if (!tenant_quota_table_.DecrementMetadataObjectCount(tenant_id)) { - LOG(WARNING) << "tenant metadata object count decrement mismatch " - << "tenant=" << tenant_id.value(); - } -} - -void MasterService::ReleaseCommittedQuotaCharge(ObjectMetadata& metadata, - uint64_t bytes) { - if (!enable_multi_tenants_ || bytes == 0) { + if (account == nullptr) { + LOG(ERROR) << "tenant quota release attempted without a bound handle" + << ", bytes=" << bytes; return; } - const uint64_t release_bytes = - std::min(bytes, metadata.committed_quota_charge_bytes); - if (release_bytes == metadata.committed_quota_charge_bytes) { - ReleaseTenantQuota(metadata.tenant_id, release_bytes); - } else { - ReleaseTenantQuotaPartial(metadata.tenant_id, release_bytes); + if (!account->Release(bytes)) { + LOG(ERROR) << "tenant quota release mismatch bytes=" << bytes; } - metadata.committed_quota_charge_bytes -= release_bytes; } void MasterService::RebuildTenantQuotaUsageFromMetadata() { @@ -1480,21 +1448,37 @@ void MasterService::RebuildTenantQuotaUsageFromMetadata() { } TenantQuotaUsageMap usage; + for (size_t i = 0; i < kNumShards; ++i) { + MetadataShardAccessorRO shard(this, i); + for (const auto& [tenant_id, tenant_state] : shard->tenants) { + for (const auto& [_, metadata] : tenant_state.metadata) { + auto& charged_bytes = usage[tenant_id]; + const uint64_t charge = CompletedMemoryQuotaCharge(metadata); + if (charge > TenantQuotaAccount::kMaxChargedBytes || + charged_bytes > + TenantQuotaAccount::kMaxChargedBytes - charge) { + throw std::overflow_error( + "rebuilt tenant quota exceeds 2^63 - 1 bytes"); + } + charged_bytes += charge; + } + } + } + for (size_t i = 0; i < kNumShards; ++i) { MetadataShardAccessorRW shard(this, i); for (auto& [tenant_id, tenant_state] : shard->tenants) { - for (auto& [_, metadata] : tenant_state.metadata) { - auto& tenant_usage = usage[tenant_id]; - ++tenant_usage.metadata_object_count; - const uint64_t charge = CompletedMemoryQuotaCharge(metadata); - metadata.reserved_quota_charge_bytes = 0; - metadata.committed_quota_charge_bytes = charge; - metadata.pending_replaced_quota_charge_bytes = 0; - if (charge == 0) { - continue; + tenant_state.quota_account = + tenant_quota_table_.GetOrCreateTenantHandle(tenant_id); + for (auto& [key, metadata] : tenant_state.metadata) { + auto rebuild_result = metadata.quota_ledger.Rebuild( + tenant_state.quota_account, + CompletedMemoryQuotaCharge(metadata)); + if (!rebuild_result) { + throw std::runtime_error( + "failed to rebuild object tenant quota ledger for " + + tenant_id.value() + "/" + key); } - tenant_usage.used_bytes += charge; - ++tenant_usage.committed_count; } } } @@ -1509,7 +1493,10 @@ void MasterService::RebuildTenantQuotaUsageFromMetadata() { } std::lock_guard recompute_lock(tenant_quota_recompute_mutex_); const uint64_t capacity = GetTenantQuotaAllocatableCapacityBytes(); - tenant_quota_table_.RebuildUsage(usage, capacity); + auto rebuild_result = tenant_quota_table_.RebuildUsage(usage, capacity); + if (!rebuild_result) { + throw std::runtime_error("failed to rebuild tenant quota usage"); + } } std::optional MasterService::GetGroupRoute( @@ -1663,6 +1650,35 @@ size_t MasterService::EraseReplicasWithCacheTotalAccounting( return erased_replicas.size(); } +tl::expected MasterService::SettlePrimaryWriteQuotaIfReady( + TenantState& tenant_state, ObjectMetadata& metadata) { + if (!enable_multi_tenants_) { + return {}; + } + if (!metadata.IsValid()) { + LOG(ERROR) << "tenant quota surviving-object settlement attempted for " + "invalid metadata, tenant=" + << metadata.tenant_id.value() + << ", key=" << metadata.user_key; + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + if (metadata.HasReplica([](const Replica& replica) { + return replica.is_memory_replica() && replica.is_processing(); + })) { + return {}; + } + + auto account = GetBoundTenantQuotaHandle(tenant_state); + auto settle_result = metadata.quota_ledger.SettlePrimaryWrite( + account, CompletedMemoryQuotaCharge(metadata)); + if (!settle_result) { + LogTenantQuotaLedgerError(settle_result, "settle_primary_write", + metadata.tenant_id, metadata.user_key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + return {}; +} + void MasterService::FinalizeRemovedReplicasAfterDurable( const OpLogEntry& durable_entry, const std::vector& replica_ids, QuotaEraseMode quota_mode) { @@ -1703,10 +1719,30 @@ void MasterService::FinalizeRemovedReplicasAfterDurable( const uint64_t erased_memory_replicas = static_cast(std::count_if( erased_replicas.begin(), erased_replicas.end(), [](const Replica& replica) { return replica.is_memory_replica(); })); - if (erased_memory_replicas > 0) { - ReleaseCommittedQuotaCharge( - metadata, SaturatingMultiply(static_cast(metadata.size), - erased_memory_replicas)); + const bool has_processing_memory = + metadata.HasReplica([](const Replica& replica) { + return replica.is_memory_replica() && replica.is_processing(); + }); + if (enable_multi_tenants_ && erased_memory_replicas > 0 && + has_processing_memory) { + const uint64_t committed_charge = + metadata.quota_ledger.CommittedBytes(); + if (metadata.size > committed_charge / erased_memory_replicas) { + LOG(ERROR) << "tenant quota removed-replica release exceeds " + "committed bytes, tenant=" + << tenant_id.value() + << ", key=" << durable_entry.object_key + << ", object_size=" << metadata.size + << ", erased_memory_replicas=" << erased_memory_replicas + << ", committed_bytes=" << committed_charge; + } else { + const uint64_t release_bytes = + metadata.size * erased_memory_replicas; + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(tenant_state), release_bytes); + LogTenantQuotaLedgerError(release_result, "release_committed", + tenant_id, durable_entry.object_key); + } } const bool erased_local_disk = std::any_of( erased_replicas.begin(), erased_replicas.end(), @@ -1722,6 +1758,12 @@ void MasterService::FinalizeRemovedReplicasAfterDurable( if (tenant_state.Empty()) { shard->tenants.erase(tenant_it); } + } else { + auto settle_result = + SettlePrimaryWriteQuotaIfReady(tenant_state, metadata); + if (settle_result && metadata.AllReplicas(&Replica::fn_is_completed)) { + tenant_state.processing_keys.erase(durable_entry.object_key); + } } } @@ -1767,8 +1809,12 @@ void MasterService::FinalizeExpiredProcessingReplicasAfterDurable( } if (!metadata.IsValid()) { accessor.Erase(); - } else if (accessor.InProcessing()) { - accessor.EraseFromProcessing(); + } else { + auto settle_result = + SettlePrimaryWriteQuotaIfReady(accessor.GetTenantState(), metadata); + if (settle_result && accessor.InProcessing()) { + accessor.EraseFromProcessing(); + } } } @@ -1804,9 +1850,9 @@ void MasterService::FinalizeExpiredReplicationTaskAfterDurable( if (!metadata.IsValid()) { accessor.Erase(); } else if (accessor.HasReplicationTask()) { - AbortTenantQuota( - tenant_id, - accessor.GetReplicationTask().reserved_quota_charge_bytes); + const auto& task = accessor.GetReplicationTask(); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); accessor.EraseReplicationTask(); } } @@ -1937,8 +1983,14 @@ MasterService::EraseMetadata( } } tenant_state.processing_keys.erase(key); - tenant_state.replication_tasks.erase(key); - ErasePromotionTaskIfPresent(tenant_state, key, tenant_id); + auto replication_task_it = tenant_state.replication_tasks.find(key); + if (replication_task_it != tenant_state.replication_tasks.end()) { + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), + replication_task_it->second.pending_quota_charge_bytes); + tenant_state.replication_tasks.erase(replication_task_it); + } + ErasePromotionTaskIfPresent(tenant_state, key); ReleaseLocalDiskUsage(metadata.GetAllReplicas()); AccountCacheTotalRemoval(metadata); @@ -1947,21 +1999,34 @@ MasterService::EraseMetadata( } switch (quota_mode) { case QuotaEraseMode::kFull: - AbortTenantQuota(tenant_id, metadata.reserved_quota_charge_bytes); - ReleaseTenantQuota(tenant_id, - metadata.committed_quota_charge_bytes); - ReleaseTenantQuota(tenant_id, - metadata.pending_replaced_quota_charge_bytes); + if (enable_multi_tenants_ && + metadata.quota_ledger.TotalChargedBytes() != 0) { + auto release_result = metadata.quota_ledger.ReleaseAll( + GetBoundTenantQuotaHandle(tenant_state)); + LogTenantQuotaLedgerError(release_result, "release_all", + tenant_id, key); + } break; case QuotaEraseMode::kPreserveOld: - AbortTenantQuota(tenant_id, metadata.reserved_quota_charge_bytes); - break; case QuotaEraseMode::kAbortOnly: - AbortTenantQuota(tenant_id, metadata.reserved_quota_charge_bytes); + if (enable_multi_tenants_ && + metadata.quota_ledger.PendingBytes() != 0) { + auto refund_result = metadata.quota_ledger.RefundPending( + GetBoundTenantQuotaHandle(tenant_state)); + LogTenantQuotaLedgerError(refund_result, "refund_pending", + tenant_id, key); + } + if (enable_multi_tenants_ && + metadata.quota_ledger.TotalChargedBytes() != 0) { + LOG(ERROR) + << "tenant quota ledger still owns bytes during preserved " + "metadata erase tenant=" + << tenant_id.value() << ", key=" << key + << ", bytes=" << metadata.quota_ledger.TotalChargedBytes(); + } break; } auto next = tenant_state.metadata.erase(it); - DecrementTenantMetadataObjectCount(tenant_id); if (had_completed_disk && shard) { shard->OnDiskReplicaRemoved(had_completed_disk); } @@ -2846,7 +2911,7 @@ void MasterService::RestoreFromStandbySnapshot( } } - auto& tenant_state = shard->tenants[tenant_id]; + auto& tenant_state = GetOrCreateTenantState(shard.get(), tenant_id); tenant_state.metadata.emplace( std::piecewise_construct, std::forward_as_tuple(user_key), std::forward_as_tuple( @@ -2862,6 +2927,10 @@ void MasterService::RestoreFromStandbySnapshot( } } + if (enable_multi_tenants_) { + RebuildTenantQuotaUsageFromMetadata(); + } + // 4. Log the result. LOG(INFO) << "Restored from standby: " << objects.size() << " objects, " << segments.size() @@ -3563,11 +3632,12 @@ auto MasterService::AllocateAndInsertMetadata( const ReplicateConfig& config, const std::string& group_id, const TenantId& tenant_id, const std::chrono::system_clock::time_point& now, const ResolvedSoftPinRequest& soft_pin_request, + uint64_t& quota_deficit_bytes, std::optional committed_soft_pin_timeout) -> tl::expected, ErrorCode> { const auto deadline_to_index = committed_soft_pin_timeout; - auto& tenant_state = shard->tenants[tenant_id]; + auto& tenant_state = GetOrCreateTenantState(shard.get(), tenant_id); if (tenant_state.metadata.contains(key)) { LOG(INFO) << "key=" << key << ", info=object_already_exists"; return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); @@ -3577,14 +3647,17 @@ auto MasterService::AllocateAndInsertMetadata( return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); } - const uint64_t reserved_quota_charge = + const uint64_t pending_quota_charge = RequestedMemoryQuotaCharge(value_length, config); - auto quota_result = ReserveTenantQuota(tenant_id, reserved_quota_charge); + auto quota_result = + ChargeTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge, "a_deficit_bytes); if (!quota_result) { return tl::make_unexpected(quota_result.error()); } - auto abort_reserved_quota = [&] { - AbortTenantQuota(tenant_id, reserved_quota_charge); + auto refund_pending_quota = [&] { + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge); }; std::vector replicas; @@ -3652,13 +3725,13 @@ auto MasterService::AllocateAndInsertMetadata( VLOG(1) << "Failed to allocate replicas for key=" << key << ", error: " << allocation_result.error(); if (allocation_result.error() == ErrorCode::INVALID_PARAMS) { - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } if (write_mode != ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA) { MasterMetricManager::instance().inc_put_start_alloc_failures(); need_mem_eviction_ = true; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } } else { @@ -3685,13 +3758,13 @@ auto MasterService::AllocateAndInsertMetadata( VLOG(1) << "Failed to allocate nof replicas for key=" << key << ", error: " << allocation_result.error(); if (allocation_result.error() == ErrorCode::INVALID_PARAMS) { - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } if (write_mode != ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA) { MasterMetricManager::instance().inc_put_start_alloc_failures(); need_nof_eviction_ = true; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } } else { @@ -3724,7 +3797,7 @@ auto MasterService::AllocateAndInsertMetadata( << ", allocated_memory_replicas=" << allocated_memory_replicas << ", requested_nof_replicas=" << config.nof_replica_num << ", allocated_nof_replicas=" << allocated_nof_replicas; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } @@ -3785,10 +3858,20 @@ auto MasterService::AllocateAndInsertMetadata( tenant_id, key)); if (!inserted) { LOG(INFO) << "key=" << key << ", info=object_already_exists"; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); } - IncrementTenantMetadataObjectCount(tenant_id); + if (enable_multi_tenants_) { + auto adopt_result = it->second.quota_ledger.AdoptPendingCharge( + GetBoundTenantQuotaHandle(tenant_state), pending_quota_charge); + if (!adopt_result) { + LogTenantQuotaLedgerError(adopt_result, "adopt_pending", tenant_id, + key); + refund_pending_quota(); + tenant_state.metadata.erase(it); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } it->second.BeginSoftPinAction(soft_pin_request, std::move(eligible_replica_ids)); if (deadline_to_index) { @@ -3796,7 +3879,6 @@ auto MasterService::AllocateAndInsertMetadata( GetMetadataShardIndex(it->second), *deadline_to_index); } - it->second.reserved_quota_charge_bytes = reserved_quota_charge; RegisterGroupMember(tenant_state, tenant_id, key, group_id); tenant_state.processing_keys.insert(key); @@ -3808,12 +3890,7 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, const uint64_t slice_length, const ReplicateConfig& config) -> tl::expected, ErrorCode> { - auto normalized_tenant_result = ResolveTenantIdForWrite(tenant_id); - if (!normalized_tenant_result) { - return tl::make_unexpected(normalized_tenant_result.error()); - } - const ObjectIdentity object_id{std::move(normalized_tenant_result.value()), - key}; + const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); if ((config.replica_num == 0 && config.nof_replica_num == 0) || key.empty() || slice_length == 0) { LOG(ERROR) << "key=" << key << ", replica_num=" << config.replica_num @@ -3865,22 +3942,11 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, [[maybe_unused]] auto object_operation_lock = AcquireObjectOperationLock(object_id.tenant_id, object_id.user_key); - const uint64_t requested_quota_charge = - RequestedMemoryQuotaCharge(slice_length, config); + uint64_t quota_deficit_bytes = 0; auto attempt_once = [&]() -> tl::expected, ErrorCode> { - std::unique_lock zero_charge_policy_lock( - tenant_quota_policy_mutex_, std::defer_lock); - if (ShouldProtectZeroChargeMetadataCreate(requested_quota_charge)) { - zero_charge_policy_lock.lock(); - auto latest_tenant_result = - ResolveTenantIdForWriteLocked(tenant_id); - if (!latest_tenant_result) { - return tl::make_unexpected(latest_tenant_result.error()); - } - } - + quota_deficit_bytes = 0; auto now = std::chrono::system_clock::now(); std::optional retry_shard_idx; { @@ -3889,7 +3955,16 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, const size_t lookup_shard_idx = getMetadataShardIndex(object_id.tenant_id, object_id.user_key); MetadataShardAccessorRW shard(this, lookup_shard_idx); - auto& tenant_state = shard->tenants[object_id.tenant_id]; + auto& tenant_state = + GetOrCreateTenantState(shard.get(), object_id.tenant_id); + auto admission_result = + ChargeTenantQuota(GetBoundTenantQuotaHandle(tenant_state), 0); + if (!admission_result) { + if (tenant_state.Empty()) { + shard->tenants.erase(object_id.tenant_id); + } + return tl::make_unexpected(admission_result.error()); + } auto it = tenant_state.metadata.find(key); if (it != tenant_state.metadata.end()) { @@ -3959,23 +4034,25 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, } else { return AllocateAndInsertMetadata( shard, client_id, key, slice_length, config, group_id, - object_id.tenant_id, now, *soft_pin_request); + object_id.tenant_id, now, *soft_pin_request, + quota_deficit_bytes); } } } std::shared_lock shared_lock(snapshot_mutex_); MetadataShardAccessorRW shard(this, retry_shard_idx.value()); - auto& retry_tenant_state = shard->tenants[object_id.tenant_id]; + auto& retry_tenant_state = + GetOrCreateTenantState(shard.get(), object_id.tenant_id); if (GetGroupRoute(object_id.tenant_id, object_id.user_key) .has_value() || retry_tenant_state.metadata.contains(key)) { LOG(INFO) << "key=" << key << ", info=object_already_exists"; return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); } - return AllocateAndInsertMetadata(shard, client_id, key, slice_length, - config, group_id, object_id.tenant_id, - now, *soft_pin_request); + return AllocateAndInsertMetadata( + shard, client_id, key, slice_length, config, group_id, + object_id.tenant_id, now, *soft_pin_request, quota_deficit_bytes); }; for (int attempt = 0; attempt <= kMaxTenantQuotaEvictionRetries; @@ -3990,10 +4067,7 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, object_id.tenant_id.value(), "quota_exceeded"); return result; } - EvictTenantMemoryForQuota( - object_id.tenant_id, - tenant_quota_table_.ComputeDeficit(object_id.tenant_id, - requested_quota_charge)); + EvictTenantMemoryForQuota(object_id.tenant_id, quota_deficit_bytes); } return tl::make_unexpected(ErrorCode::TENANT_QUOTA_EXCEEDED); } @@ -4088,28 +4162,10 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, metadata.object_checksum = object_meta.object_checksum; } - const bool has_memory_replica = metadata.HasMemReplica(); - const bool should_settle_quota = - replica_type == ReplicaType::MEMORY || - (replica_type == ReplicaType::ALL && has_memory_replica) || - !has_memory_replica; - if (metadata.reserved_quota_charge_bytes > 0 && should_settle_quota) { - const uint64_t actual_charge = CompletedMemoryQuotaCharge(metadata); - const uint64_t commit_charge = - actual_charge > metadata.committed_quota_charge_bytes - ? actual_charge - metadata.committed_quota_charge_bytes - : 0; - const uint64_t abort_charge = - metadata.reserved_quota_charge_bytes > commit_charge - ? metadata.reserved_quota_charge_bytes - commit_charge - : 0; - CommitTenantQuota(object_id.tenant_id, commit_charge); - AbortTenantQuota(object_id.tenant_id, abort_charge); - metadata.reserved_quota_charge_bytes = 0; - metadata.committed_quota_charge_bytes = actual_charge; - ReleaseTenantQuota(object_id.tenant_id, - metadata.pending_replaced_quota_charge_bytes); - metadata.pending_replaced_quota_charge_bytes = 0; + auto settle_result = + SettlePrimaryWriteQuotaIfReady(accessor.GetTenantState(), metadata); + if (!settle_result) { + return tl::make_unexpected(settle_result.error()); } if (enable_offload_ && !offload_on_evict_) { @@ -4356,13 +4412,25 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, EraseReplicasWithCacheTotalAccounting(metadata, target_pred); metadata.ClearPendingSoftPinIfNoViableReplica(); const uint64_t after_charge = CompletedMemoryQuotaCharge(metadata); - if (before_charge > after_charge) { - ReleaseCommittedQuotaCharge(metadata, before_charge - after_charge); + if (enable_multi_tenants_ && before_charge > after_charge) { + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(accessor.GetTenantState()), + before_charge - after_charge); + if (!release_result) { + LogTenantQuotaLedgerError(release_result, "release_committed", + object_id.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } + if (!metadata.IsValid()) { + accessor.Erase(); + return {}; } - if (!metadata.HasReplica(&Replica::fn_is_memory_replica)) { - AbortTenantQuota(object_id.tenant_id, - metadata.reserved_quota_charge_bytes); - metadata.reserved_quota_charge_bytes = 0; + + auto settle_result = + SettlePrimaryWriteQuotaIfReady(accessor.GetTenantState(), metadata); + if (!settle_result) { + return tl::make_unexpected(settle_result.error()); } // If the object is completed, remove it from the processing set. @@ -4371,9 +4439,6 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, accessor.EraseFromProcessing(); } - if (metadata.IsValid() == false) { - accessor.Erase(); - } return {}; } @@ -4429,12 +4494,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, const uint64_t slice_length, const ReplicateConfig& config) -> tl::expected, ErrorCode> { - auto normalized_tenant_result = ResolveTenantIdForWrite(tenant_id); - if (!normalized_tenant_result) { - return tl::make_unexpected(normalized_tenant_result.error()); - } - const ObjectIdentity object_id{std::move(normalized_tenant_result.value()), - key}; + const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); // --- Parameter validation (same as PutStart) --- if ((config.replica_num == 0 && config.nof_replica_num == 0) || key.empty() || slice_length == 0) { @@ -4487,22 +4547,11 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, [[maybe_unused]] auto object_operation_lock = AcquireObjectOperationLock(object_id.tenant_id, object_id.user_key); - const uint64_t requested_quota_charge = - RequestedMemoryQuotaCharge(slice_length, config); + uint64_t quota_deficit_bytes = 0; auto attempt_once = [&]() -> tl::expected, ErrorCode> { - std::unique_lock zero_charge_policy_lock( - tenant_quota_policy_mutex_, std::defer_lock); - if (ShouldProtectZeroChargeMetadataCreate(requested_quota_charge)) { - zero_charge_policy_lock.lock(); - auto latest_tenant_result = - ResolveTenantIdForWriteLocked(tenant_id); - if (!latest_tenant_result) { - return tl::make_unexpected(latest_tenant_result.error()); - } - } - + quota_deficit_bytes = 0; auto now = std::chrono::system_clock::now(); std::optional case_a_retry_shard_idx; std::optional @@ -4516,7 +4565,16 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, const size_t lookup_shard_idx = getMetadataShardIndex(object_id.tenant_id, object_id.user_key); MetadataShardAccessorRW shard(this, lookup_shard_idx); - auto& tenant_state = shard->tenants[object_id.tenant_id]; + auto& tenant_state = + GetOrCreateTenantState(shard.get(), object_id.tenant_id); + auto admission_result = + ChargeTenantQuota(GetBoundTenantQuotaHandle(tenant_state), 0); + if (!admission_result) { + if (tenant_state.Empty()) { + shard->tenants.erase(object_id.tenant_id); + } + return tl::make_unexpected(admission_result.error()); + } auto it = tenant_state.metadata.find(key); @@ -4614,6 +4672,12 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, EraseMetadata(tenant_state, it, object_id.tenant_id, QuotaEraseMode::kFull, &shard); it = tenant_state.metadata.end(); + } else { + auto settle_result = SettlePrimaryWriteQuotaIfReady( + tenant_state, metadata); + if (!settle_result) { + return tl::make_unexpected(settle_result.error()); + } } } } @@ -4635,6 +4699,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, return AllocateAndInsertMetadata( shard, client_id, key, slice_length, config, group_id, object_id.tenant_id, now, *soft_pin_request, + quota_deficit_bytes, std::move(case_a_committed_soft_pin_timeout)); } } else { @@ -4714,10 +4779,22 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, } const std::string existing_group_id = metadata.group_id; - const uint64_t old_quota_charge = - metadata.committed_quota_charge_bytes != 0 - ? metadata.committed_quota_charge_bytes - : CompletedMemoryQuotaCharge(metadata); + TenantQuotaLedger replacement_charge; + auto* quota_account = GetBoundTenantQuotaHandle(tenant_state); + const bool has_replacement_charge = + enable_multi_tenants_ && + metadata.quota_ledger.TotalChargedBytes() != 0; + if (has_replacement_charge) { + auto transfer_result = + metadata.quota_ledger.TransferReplacementCharge( + quota_account, replacement_charge); + if (!transfer_result) { + LogTenantQuotaLedgerError(transfer_result, + "transfer_replacement_out", + object_id.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } auto old_replicas = PopReplicasWithCacheTotalAccounting(metadata); if (!old_replicas.empty()) { @@ -4734,22 +4811,57 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, auto allocate_result = AllocateAndInsertMetadata( shard, client_id, key, slice_length, merged_config, existing_group_id, object_id.tenant_id, now, - *soft_pin_request, std::move(committed_soft_pin_timeout)); + *soft_pin_request, quota_deficit_bytes, + std::move(committed_soft_pin_timeout)); if (!allocate_result) { - ReleaseTenantQuota(object_id.tenant_id, old_quota_charge); + if (has_replacement_charge) { + auto rollback_result = + replacement_charge.ReleaseReplacement( + quota_account); + LogTenantQuotaLedgerError(rollback_result, + "rollback_replacement", + object_id.tenant_id, key); + } return allocate_result; } auto new_it = tenant_state.metadata.find(key); - if (new_it != tenant_state.metadata.end()) { - new_it->second.pending_replaced_quota_charge_bytes = - old_quota_charge; + if (has_replacement_charge) { + if (new_it == tenant_state.metadata.end()) { + auto rollback_result = + replacement_charge.ReleaseReplacement( + quota_account); + LogTenantQuotaLedgerError(rollback_result, + "rollback_replacement", + object_id.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + auto transfer_result = + replacement_charge.TransferReplacementCharge( + quota_account, new_it->second.quota_ledger); + if (!transfer_result) { + LogTenantQuotaLedgerError(transfer_result, + "transfer_replacement_in", + object_id.tenant_id, key); + EraseMetadata(tenant_state, new_it, object_id.tenant_id, + QuotaEraseMode::kFull, &shard); + if (replacement_charge.ReplacedBytes() != 0) { + auto rollback_result = + replacement_charge.ReleaseReplacement( + quota_account); + LogTenantQuotaLedgerError(rollback_result, + "rollback_replacement", + object_id.tenant_id, key); + } + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } } return allocate_result; } } std::shared_lock shared_lock(snapshot_mutex_); MetadataShardAccessorRW shard(this, case_a_retry_shard_idx.value()); - auto& retry_tenant_state = shard->tenants[object_id.tenant_id]; + auto& retry_tenant_state = + GetOrCreateTenantState(shard.get(), object_id.tenant_id); const auto current_route = GetGroupRoute(object_id.tenant_id, object_id.user_key); if (current_route.has_value() || @@ -4759,7 +4871,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, } return AllocateAndInsertMetadata( shard, client_id, key, slice_length, config, group_id, - object_id.tenant_id, now, *soft_pin_request, + object_id.tenant_id, now, *soft_pin_request, quota_deficit_bytes, std::move(case_a_committed_soft_pin_timeout)); }; @@ -4775,10 +4887,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, object_id.tenant_id.value(), "quota_exceeded"); return result; } - EvictTenantMemoryForQuota( - object_id.tenant_id, - tenant_quota_table_.ComputeDeficit(object_id.tenant_id, - requested_quota_charge)); + EvictTenantMemoryForQuota(object_id.tenant_id, quota_deficit_bytes); } return tl::make_unexpected(ErrorCode::TENANT_QUOTA_EXCEEDED); } @@ -4984,12 +5093,7 @@ tl::expected MasterService::CopyStart( const UUID& client_id, const std::string& key, const TenantId& tenant_id, const std::string& src_segment, const std::vector& tgt_segments) { - auto normalized_tenant_result = ResolveTenantIdForWrite(tenant_id); - if (!normalized_tenant_result) { - return tl::make_unexpected(normalized_tenant_result.error()); - } - const ObjectIdentity object_id{std::move(normalized_tenant_result.value()), - key}; + const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); std::shared_lock shared_lock(snapshot_mutex_); { ScopedSegmentAccess segment_access = @@ -5021,6 +5125,7 @@ tl::expected MasterService::CopyStart( } auto& metadata = accessor.Get(); + auto& tenant_state = accessor.GetTenantState(); auto source = metadata.GetReplicaBySegmentName(src_segment); if (source == nullptr || !source->is_completed() || source->has_invalid_mem_handle()) { @@ -5036,11 +5141,11 @@ tl::expected MasterService::CopyStart( } } - const uint64_t reserved_quota_charge = + const uint64_t pending_quota_charge = SaturatingMultiply(static_cast(metadata.size), static_cast(new_replica_count)); - auto quota_result = - ReserveTenantQuota(object_id.tenant_id, reserved_quota_charge); + auto quota_result = ChargeTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), pending_quota_charge); if (!quota_result) { if (quota_result.error() == ErrorCode::TENANT_QUOTA_EXCEEDED) { MasterMetricManager::instance().inc_tenant_quota_reject( @@ -5048,8 +5153,9 @@ tl::expected MasterService::CopyStart( } return tl::make_unexpected(quota_result.error()); } - auto abort_reserved_quota = [&] { - AbortTenantQuota(object_id.tenant_id, reserved_quota_charge); + auto refund_pending_quota = [&] { + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge); }; std::vector replicas; @@ -5070,7 +5176,7 @@ tl::expected MasterService::CopyStart( if (!replica.has_value()) { LOG(ERROR) << "key=" << key << ", tgt_segment=" << tgt_segment << ", failed to allocate replica"; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(replica.error()); } replicas.push_back(std::move(*replica)); @@ -5089,14 +5195,13 @@ tl::expected MasterService::CopyStart( } // Create replication task for tracking. - auto& tenant_state = accessor.GetTenantState(); auto task_insert = tenant_state.replication_tasks.emplace( std::piecewise_construct, std::forward_as_tuple(key), std::forward_as_tuple(client_id, std::chrono::system_clock::now(), ReplicationTask::Type::COPY, source->id(), - std::move(replica_ids), reserved_quota_charge)); + std::move(replica_ids), pending_quota_charge)); if (!task_insert.second) { - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::OBJECT_HAS_REPLICATION_TASK); } @@ -5159,7 +5264,8 @@ tl::expected MasterService::CopyEnd( task.replica_ids.end(), replica.id()) != task.replica_ids.end(); }); - AbortTenantQuota(metadata.tenant_id, task.reserved_quota_charge_bytes); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); accessor.EraseReplicationTask(); if (!metadata.IsValid()) { // Remove the object if it does not have any replicas. @@ -5237,14 +5343,16 @@ tl::expected MasterService::CopyEnd( SyncCacheTotalAccounting(metadata); - const uint64_t commit_charge = - std::min(completed_quota_charge, task.reserved_quota_charge_bytes); - const uint64_t abort_charge = - task.reserved_quota_charge_bytes - commit_charge; - CommitAdditionalTenantQuota(metadata.tenant_id, commit_charge); - AbortTenantQuota(metadata.tenant_id, abort_charge); - metadata.committed_quota_charge_bytes = - SaturatingAdd(metadata.committed_quota_charge_bytes, commit_charge); + if (enable_multi_tenants_) { + auto settle_result = metadata.quota_ledger.SettleAdditional( + GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes, completed_quota_charge); + if (!settle_result) { + LogTenantQuotaLedgerError(settle_result, "settle_additional", + metadata.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } accessor.EraseReplicationTask(); @@ -5299,7 +5407,8 @@ tl::expected MasterService::CopyRevoke( }); } - AbortTenantQuota(metadata.tenant_id, task.reserved_quota_charge_bytes); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); accessor.EraseReplicationTask(); if (!metadata.IsValid()) { @@ -5313,12 +5422,7 @@ tl::expected MasterService::CopyRevoke( tl::expected MasterService::MoveStart( const UUID& client_id, const std::string& key, const TenantId& tenant_id, const std::string& src_segment, const std::string& tgt_segment) { - auto normalized_tenant_result = ResolveTenantIdForWrite(tenant_id); - if (!normalized_tenant_result) { - return tl::make_unexpected(normalized_tenant_result.error()); - } - const ObjectIdentity object_id{std::move(normalized_tenant_result.value()), - key}; + const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); std::shared_lock shared_lock(snapshot_mutex_); if (src_segment == tgt_segment) { LOG(ERROR) << "key=" << key << ", move_tgt=" << tgt_segment @@ -5354,6 +5458,7 @@ tl::expected MasterService::MoveStart( } auto& metadata = accessor.Get(); + auto& tenant_state = accessor.GetTenantState(); auto source = metadata.GetReplicaBySegmentName(src_segment); if (source == nullptr || !source->is_completed() || source->has_invalid_mem_handle()) { @@ -5364,10 +5469,10 @@ tl::expected MasterService::MoveStart( std::vector replicas; if (metadata.GetReplicaBySegmentName(tgt_segment) == nullptr) { - const uint64_t reserved_quota_charge = + const uint64_t pending_quota_charge = SaturatingMultiply(static_cast(metadata.size), 1); - auto quota_result = - ReserveTenantQuota(object_id.tenant_id, reserved_quota_charge); + auto quota_result = ChargeTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), pending_quota_charge); if (!quota_result) { if (quota_result.error() == ErrorCode::TENANT_QUOTA_EXCEEDED) { MasterMetricManager::instance().inc_tenant_quota_reject( @@ -5375,8 +5480,9 @@ tl::expected MasterService::MoveStart( } return tl::make_unexpected(quota_result.error()); } - auto abort_reserved_quota = [&] { - AbortTenantQuota(object_id.tenant_id, reserved_quota_charge); + auto refund_pending_quota = [&] { + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge); }; ScopedAllocatorAccess allocator_access = @@ -5388,18 +5494,19 @@ tl::expected MasterService::MoveStart( if (!replica.has_value()) { LOG(ERROR) << "key=" << key << ", tgt_segment=" << tgt_segment << ", failed to allocate replica"; - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(replica.error()); } replicas.push_back(std::move(*replica)); } else { - auto quota_result = ReserveTenantQuota(object_id.tenant_id, 0); + auto quota_result = + ChargeTenantQuota(GetBoundTenantQuotaHandle(tenant_state), 0); if (!quota_result) { return tl::make_unexpected(quota_result.error()); } } - const uint64_t reserved_quota_charge = + const uint64_t pending_quota_charge = replicas.empty() ? 0 : SaturatingMultiply(static_cast(metadata.size), 1); @@ -5416,14 +5523,14 @@ tl::expected MasterService::MoveStart( } // Create replication task for tracking. - auto& tenant_state = accessor.GetTenantState(); auto task_insert = tenant_state.replication_tasks.emplace( std::piecewise_construct, std::forward_as_tuple(key), std::forward_as_tuple(client_id, std::chrono::system_clock::now(), ReplicationTask::Type::MOVE, source->id(), - std::move(replica_ids), reserved_quota_charge)); + std::move(replica_ids), pending_quota_charge)); if (!task_insert.second) { - AbortTenantQuota(object_id.tenant_id, reserved_quota_charge); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge); return tl::make_unexpected(ErrorCode::OBJECT_HAS_REPLICATION_TASK); } @@ -5486,7 +5593,8 @@ tl::expected MasterService::MoveEnd( task.replica_ids.end(), replica.id()) != task.replica_ids.end(); }); - AbortTenantQuota(metadata.tenant_id, task.reserved_quota_charge_bytes); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); accessor.EraseReplicationTask(); if (!metadata.IsValid()) { // Remove the object if it does not have any replicas. @@ -5505,8 +5613,9 @@ tl::expected MasterService::MoveEnd( LOG(WARNING) << "key=" << key << ", replica_id=" << target_id << ", move target becomes invalid during data transfer"; - AbortTenantQuota(metadata.tenant_id, - task.reserved_quota_charge_bytes); + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); // Source untouched; safe to drop the broken task. accessor.EraseReplicationTask(); return tl::make_unexpected(ErrorCode::REPLICA_IS_GONE); @@ -5571,6 +5680,17 @@ tl::expected MasterService::MoveEnd( replica->mark_complete(); } } + if (enable_multi_tenants_) { + auto settle_result = metadata.quota_ledger.SettleAdditional( + GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes, + has_target ? static_cast(metadata.size) : 0); + if (!settle_result) { + LogTenantQuotaLedgerError(settle_result, "settle_additional", + metadata.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } if (!(enable_ha_ && enable_oplog_)) { // Remove the source replica and release its space later. @@ -5579,6 +5699,17 @@ tl::expected MasterService::MoveEnd( return replica.id() == source_id; }); if (!source_replica.empty()) { + if (enable_multi_tenants_) { + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(accessor.GetTenantState()), + static_cast(metadata.size)); + if (!release_result) { + LogTenantQuotaLedgerError(release_result, + "release_committed", + metadata.tenant_id, key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back( std::move(source_replica), std::chrono::system_clock::now() + @@ -5586,7 +5717,6 @@ tl::expected MasterService::MoveEnd( } } - AbortTenantQuota(metadata.tenant_id, task.reserved_quota_charge_bytes); accessor.EraseReplicationTask(); return {}; @@ -5639,7 +5769,8 @@ tl::expected MasterService::MoveRevoke( }); } - AbortTenantQuota(metadata.tenant_id, task.reserved_quota_charge_bytes); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(accessor.GetTenantState()), + task.pending_quota_charge_bytes); accessor.EraseReplicationTask(); if (!metadata.IsValid()) { @@ -5895,8 +6026,7 @@ long MasterService::RemoveAll(bool force) { } total_freed_size += it->second.size * mem_rep_count; - ErasePromotionTaskIfPresent(tenant_state, it->first, - tenant_it->first); + ErasePromotionTaskIfPresent(tenant_state, it->first); it = EraseMetadata(tenant_state, it, tenant_it->first, QuotaEraseMode::kFull, &shard); removed_count++; @@ -5987,8 +6117,7 @@ long MasterService::RemoveAll(const TenantId& tenant_id, bool force) { } } total_freed_size += it->second.size * mem_rep_count; - ErasePromotionTaskIfPresent(tenant_state, it->first, - normalized_tenant); + ErasePromotionTaskIfPresent(tenant_state, it->first); it = EraseMetadata(tenant_state, it, normalized_tenant, QuotaEraseMode::kFull, &shard); removed_count++; @@ -6196,8 +6325,7 @@ void MasterService::CancelPromotionTaskForRemovedReplicas( source->dec_refcnt(); } const UUID holder_id = task_it->second.holder_id; - ErasePromotionTaskIfPresent(tenant_state, metadata.user_key, - metadata.tenant_id); + ErasePromotionTaskIfPresent(tenant_state, metadata.user_key); // Best-effort cleanup of a task that may still be queued on the holder. ScopedLocalDiskSegmentAccess local_disk_segment_access = @@ -6235,8 +6363,12 @@ bool MasterService::CleanupStaleHandles( CancelPromotionTaskForRemovedReplicas(tenant_state, metadata, removed_replica_ids); const uint64_t after_charge = CompletedMemoryQuotaCharge(metadata); - if (before_charge > after_charge) { - ReleaseCommittedQuotaCharge(metadata, before_charge - after_charge); + if (enable_multi_tenants_ && before_charge > after_charge) { + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(tenant_state), + before_charge - after_charge); + LogTenantQuotaLedgerError(release_result, "release_committed", + metadata.tenant_id, metadata.user_key); } if (had_completed_disk && shard && !metadata.HasReplica([](const Replica& r) { @@ -7170,12 +7302,7 @@ auto MasterService::PromotionAllocStart( const UUID& client_id, const std::string& key, const TenantId& tenant_id, uint64_t size, const std::vector& preferred_segments) -> tl::expected { - auto normalized_tenant_result = ResolveTenantIdForWrite(tenant_id); - if (!normalized_tenant_result) { - return tl::make_unexpected(normalized_tenant_result.error()); - } - const ObjectIdentity object_id{std::move(normalized_tenant_result.value()), - key}; + const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); std::shared_lock shared_lock(snapshot_mutex_); MetadataAccessorRW accessor(this, object_id); if (!accessor.Exists()) { @@ -7220,18 +7347,19 @@ auto MasterService::PromotionAllocStart( return tl::make_unexpected(ErrorCode::REPLICA_IS_NOT_READY); } if (task_it->second.alloc_id != 0 || - task_it->second.reserved_quota_charge_bytes != 0) { + task_it->second.pending_quota_charge_bytes != 0) { return tl::make_unexpected(ErrorCode::REPLICA_IS_NOT_READY); } - const uint64_t reserved_quota_charge = size; - auto quota_result = - ReserveTenantQuota(object_id.tenant_id, reserved_quota_charge); + const uint64_t pending_quota_charge = size; + auto quota_result = ChargeTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), pending_quota_charge); if (!quota_result) { return tl::make_unexpected(quota_result.error()); } - auto abort_reserved_quota = [&] { - AbortTenantQuota(object_id.tenant_id, reserved_quota_charge); + auto refund_pending_quota = [&] { + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + pending_quota_charge); }; // Allocate a single MEMORY replica via the existing strategy, biased to @@ -7250,13 +7378,13 @@ auto MasterService::PromotionAllocStart( auto allocation_result = allocation_strategy_->Allocate( allocator_manager, size, config.replica_num, preferred_segments); if (!allocation_result) { - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(allocation_result.error()); } staged_replicas = std::move(allocation_result.value()); } if (staged_replicas.empty()) { - abort_reserved_quota(); + refund_pending_quota(); return tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } @@ -7284,7 +7412,7 @@ auto MasterService::PromotionAllocStart( // (alloc_id == 0) is bounded by its own original start_time window // during which the reaper's EraseReplicaByID branch is a no-op. task_it->second.alloc_id = new_id; - task_it->second.reserved_quota_charge_bytes = reserved_quota_charge; + task_it->second.pending_quota_charge_bytes = pending_quota_charge; task_it->second.start_time = std::chrono::system_clock::now(); return PromotionAllocStartResponse{std::move(desc)}; } @@ -7372,25 +7500,23 @@ auto MasterService::NotifyPromotionSuccess(const UUID& client_id, source->dec_refcnt(); } const uint64_t completed_bytes = task_it->second.object_size; - const uint64_t reserved_quota_charge = - task_it->second.reserved_quota_charge_bytes; if (committed) { - const uint64_t actual_charge = CompletedMemoryQuotaCharge(metadata); - const uint64_t commit_charge = - actual_charge > metadata.committed_quota_charge_bytes - ? actual_charge - metadata.committed_quota_charge_bytes - : 0; - const uint64_t abort_charge = - reserved_quota_charge > commit_charge - ? reserved_quota_charge - commit_charge - : 0; - CommitTenantQuota(object_id.tenant_id, commit_charge); - AbortTenantQuota(object_id.tenant_id, abort_charge); - metadata.committed_quota_charge_bytes = actual_charge; + if (enable_multi_tenants_) { + auto settle_result = metadata.quota_ledger.SettleAdditional( + GetBoundTenantQuotaHandle(tenant_state), + task_it->second.pending_quota_charge_bytes, completed_bytes); + if (!settle_result) { + LogTenantQuotaLedgerError(settle_result, "settle_additional", + object_id.tenant_id, + object_id.user_key); + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + } } else { - AbortTenantQuota(object_id.tenant_id, reserved_quota_charge); + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), + std::exchange(task_it->second.pending_quota_charge_bytes, 0)); } - task_it->second.reserved_quota_charge_bytes = 0; tenant_state.promotion_tasks.erase(task_it); promotion_in_flight_.fetch_sub(1, std::memory_order_relaxed); MasterMetricManager::instance().dec_promotion_in_flight(); @@ -7464,9 +7590,9 @@ auto MasterService::NotifyPromotionFailure(const UUID& client_id, return replica.id() == alloc_id; }); } - AbortTenantQuota(object_id.tenant_id, - task_it->second.reserved_quota_charge_bytes); - task_it->second.reserved_quota_charge_bytes = 0; + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), + std::exchange(task_it->second.pending_quota_charge_bytes, 0)); tenant_state.promotion_tasks.erase(task_it); promotion_in_flight_.fetch_sub(1, std::memory_order_relaxed); MasterMetricManager::instance().dec_promotion_in_flight(); @@ -7585,7 +7711,13 @@ void MasterService::DiscardExpiredProcessingReplicas( QuotaEraseMode::kFull, &shard); key_it = next_key_it; } else { - key_it = tenant_state.processing_keys.erase(key_it); + auto settle_result = + SettlePrimaryWriteQuotaIfReady(tenant_state, metadata); + if (!settle_result) { + ++key_it; + } else { + key_it = tenant_state.processing_keys.erase(key_it); + } } continue; } @@ -7656,7 +7788,13 @@ void MasterService::DiscardExpiredProcessingReplicas( QuotaEraseMode::kFull, &shard); key_it = next_key_it; } else { - key_it = tenant_state.processing_keys.erase(key_it); + auto settle_result = + SettlePrimaryWriteQuotaIfReady(tenant_state, metadata); + if (!settle_result) { + ++key_it; + } else { + key_it = tenant_state.processing_keys.erase(key_it); + } } continue; } @@ -7669,8 +7807,8 @@ void MasterService::DiscardExpiredProcessingReplicas( if (metadata_it == tenant_state.metadata.end()) { LOG(ERROR) << "Key " << task_it->first << " was removed with ongoing replication task"; - AbortTenantQuota(tenant_it->first, - task_it->second.reserved_quota_charge_bytes); + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + task_it->second.pending_quota_charge_bytes); task_it = tenant_state.replication_tasks.erase(task_it); continue; } @@ -7765,6 +7903,8 @@ void MasterService::DiscardExpiredProcessingReplicas( QuotaEraseMode::kFull, &shard); task_it = next_task_it; } else { + ReleaseTenantQuota(GetBoundTenantQuotaHandle(tenant_state), + task_it->second.pending_quota_charge_bytes); task_it = tenant_state.replication_tasks.erase(task_it); } } @@ -7814,9 +7954,9 @@ void MasterService::DiscardExpiredProcessingReplicas( }); } } - AbortTenantQuota(tenant_it->first, - task_it->second.reserved_quota_charge_bytes); - task_it->second.reserved_quota_charge_bytes = 0; + ReleaseTenantQuota( + GetBoundTenantQuotaHandle(tenant_state), + std::exchange(task_it->second.pending_quota_charge_bytes, 0)); LOG(WARNING) << "Promotion task expired for key: " << task_it->first; task_it = tenant_state.promotion_tasks.erase(task_it); @@ -8124,7 +8264,7 @@ MasterService::EvictTenantMemoryForQuota(const TenantId& tenant_id, return metadata.HasReplica(&Replica::fn_is_local_disk_replica); }; auto evict_replicas = - [&, this](ObjectMetadata& metadata, + [&, this](TenantState& tenant_state, ObjectMetadata& metadata, std::vector>& deferred_replicas) { const uint64_t before_charge = CompletedMemoryQuotaCharge(metadata); auto replicas = PopReplicasWithCacheTotalAccounting( @@ -8135,8 +8275,12 @@ MasterService::EvictTenantMemoryForQuota(const TenantId& tenant_id, } const uint64_t after_charge = CompletedMemoryQuotaCharge(metadata); if (before_charge > after_charge) { - ReleaseCommittedQuotaCharge(metadata, - before_charge - after_charge); + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(tenant_state), + before_charge - after_charge); + LogTenantQuotaLedgerError(release_result, "release_committed", + metadata.tenant_id, + metadata.user_key); } return metadata.size * replica_count; }; @@ -8149,54 +8293,54 @@ MasterService::EvictTenantMemoryForQuota(const TenantId& tenant_id, ? static_cast(offloading_queue_limit_ * offload_cap_ratio_) : 0; - auto try_evict_or_offload = - [&, this](const std::string& key, ObjectMetadata& metadata, - TenantState& tenant_state, - std::vector>& deferred_replicas) { - if (!offload_on_evict_) { - return evict_replicas(metadata, deferred_replicas); - } + auto try_evict_or_offload = [&, this](const std::string& key, + ObjectMetadata& metadata, + TenantState& tenant_state, + std::vector>& + deferred_replicas) { + if (!offload_on_evict_) { + return evict_replicas(tenant_state, metadata, deferred_replicas); + } - if (has_local_disk_replica(metadata)) { - return evict_replicas(metadata, deferred_replicas); - } + if (has_local_disk_replica(metadata)) { + return evict_replicas(tenant_state, metadata, deferred_replicas); + } - if (offload_force_evict_ && - offload_queued_this_call >= offload_cap) { - ++offload_cap_forced_count; - return evict_replicas(metadata, deferred_replicas); - } + if (offload_force_evict_ && offload_queued_this_call >= offload_cap) { + ++offload_cap_forced_count; + return evict_replicas(tenant_state, metadata, deferred_replicas); + } - bool queued = false; - metadata.VisitReplicas( - is_evictable_memory_replica, - [this, &key, &normalized_tenant, &tenant_state, &queued, - &now](Replica& replica) { - if (queued) { - return; - } - auto result = PushOffloadingQueue( - MakeObjectIdentity(key, normalized_tenant), replica); - if (result) { - replica.inc_refcnt(); - tenant_state.offloading_tasks.emplace( - key, OffloadingTask{replica.id(), now}); - queued = true; - } - }); + bool queued = false; + metadata.VisitReplicas( + is_evictable_memory_replica, + [this, &key, &normalized_tenant, &tenant_state, &queued, + &now](Replica& replica) { + if (queued) { + return; + } + auto result = PushOffloadingQueue( + MakeObjectIdentity(key, normalized_tenant), replica); + if (result) { + replica.inc_refcnt(); + tenant_state.offloading_tasks.emplace( + key, OffloadingTask{replica.id(), now}); + queued = true; + } + }); - if (queued) { - ++offload_queued_this_call; - ++offload_deferred_count; - return evict_replicas(metadata, deferred_replicas); - } + if (queued) { + ++offload_queued_this_call; + ++offload_deferred_count; + return evict_replicas(tenant_state, metadata, deferred_replicas); + } - if (offload_force_evict_) { - ++offload_push_failed_forced; - return evict_replicas(metadata, deferred_replicas); - } - return uint64_t{0}; - }; + if (offload_force_evict_) { + ++offload_push_failed_forced; + return evict_replicas(tenant_state, metadata, deferred_replicas); + } + return uint64_t{0}; + }; auto try_evict_group_or_object = [&, this](const std::string& key, ObjectMetadata& metadata, @@ -8355,7 +8499,7 @@ void MasterService::BatchEvict(double evict_ratio_target, }; auto evict_replicas = - [&, this](ObjectMetadata& metadata, + [&, this](TenantState& tenant_state, ObjectMetadata& metadata, std::vector>& deferred_replicas) { if (enable_oplog_) { return metadata.size * @@ -8372,9 +8516,13 @@ void MasterService::BatchEvict(double evict_ratio_target, deferred_replicas.emplace_back(std::move(replicas)); } const uint64_t after_charge = CompletedMemoryQuotaCharge(metadata); - if (before_charge > after_charge) { - ReleaseCommittedQuotaCharge(metadata, - before_charge - after_charge); + if (enable_multi_tenants_ && before_charge > after_charge) { + auto release_result = metadata.quota_ledger.ReleaseCommitted( + GetBoundTenantQuotaHandle(tenant_state), + before_charge - after_charge); + LogTenantQuotaLedgerError(release_result, "release_committed", + metadata.tenant_id, + metadata.user_key); } return metadata.size * replica_count; }; @@ -8401,16 +8549,16 @@ void MasterService::BatchEvict(double evict_ratio_target, ObjectMetadata& metadata, TenantState& tenant_state, std::vector>& deferred_replicas) -> uint64_t { if (enable_oplog_) { - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } if (!offload_on_evict_) { // Original behavior - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } // LOCAL_DISK replica already exists — safe to delete MEMORY immediately if (has_local_disk_replica(metadata)) { - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } // Force-evict cap: if force_evict enabled and cap reached, force @@ -8418,7 +8566,7 @@ void MasterService::BatchEvict(double evict_ratio_target, // flooding. if (offload_force_evict_ && offload_queued_this_cycle >= offload_cap) { offload_cap_forced_count++; - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } // Queue one MEMORY replica for offload; others will be evicted below. @@ -8447,7 +8595,7 @@ void MasterService::BatchEvict(double evict_ratio_target, // Any remaining MEMORY replicas with refcnt==0 are redundant copies // (data survives via the pinned replica → disk). Evict them now to // reclaim memory immediately rather than waiting another cycle. - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } // PushOffloadingQueue failed. Default (data-preserving) behavior is to @@ -8456,7 +8604,7 @@ void MasterService::BatchEvict(double evict_ratio_target, // prevent silent data loss when the queue is unavailable. if (offload_force_evict_) { offload_push_failed_forced++; - return evict_replicas(metadata, deferred_replicas); + return evict_replicas(tenant_state, metadata, deferred_replicas); } return 0; }; @@ -10081,7 +10229,7 @@ MasterService::MetadataSerializer::DeserializeShard(const msgpack::object& obj, } auto metadata_ptr = std::move(metadata_result.value()); - auto& tenant_state = shard.tenants[tenant_id]; + auto& tenant_state = service_->GetOrCreateTenantState(shard, tenant_id); const std::string user_key = key; auto [it, inserted] = tenant_state.metadata.emplace( std::piecewise_construct, std::forward_as_tuple(std::move(key)), diff --git a/mooncake-store/src/tenant_quota.cpp b/mooncake-store/src/tenant_quota.cpp index fbd7372927..14c73ca176 100644 --- a/mooncake-store/src/tenant_quota.cpp +++ b/mooncake-store/src/tenant_quota.cpp @@ -1,6 +1,7 @@ #include "tenant_quota.h" #include +#include #include namespace mooncake { @@ -12,11 +13,16 @@ struct RemainderShare { unsigned __int128 remainder = 0; }; -uint64_t SaturatingAdd(uint64_t lhs, uint64_t rhs) { - if (lhs > std::numeric_limits::max() - rhs) { - return std::numeric_limits::max(); - } - return lhs + rhs; +TenantQuotaResult AccountingMismatch() { + return tl::make_unexpected(TenantQuotaError::kAccountingMismatch); +} + +TenantQuotaChargeResult ChargeFailure(TenantQuotaError error, + uint64_t deficit_bytes = 0) { + return tl::make_unexpected(TenantQuotaChargeFailure{ + .error = error, + .deficit_bytes = deficit_bytes, + }); } std::map BuildEffectiveQuotaAssignmentsImpl( @@ -78,11 +84,150 @@ std::map BuildEffectiveQuotaAssignmentsImpl( return assigned; } -TenantQuotaResult AccountingMismatch() { - return tl::make_unexpected(TenantQuotaError::kAccountingMismatch); +} // namespace + +TenantQuotaChargeResult TenantQuotaAccount::TryCharge(uint64_t bytes) { + if (bytes > kMaxChargedBytes) { + return ChargeFailure(TenantQuotaError::kInvalidArgument); + } + + for (;;) { + const uint64_t sequence_before = + policy_sequence_.load(std::memory_order_acquire); + if (sequence_before & 1) { + continue; + } + + uint64_t expected = charged_state_.load(std::memory_order_acquire); + if (expected & kAdmissionClosed) { + const uint64_t sequence_after = + policy_sequence_.load(std::memory_order_acquire); + if (sequence_before == sequence_after && !(sequence_after & 1)) { + return ChargeFailure(TenantQuotaError::kTenantNotRegistered); + } + continue; + } + + const uint64_t charged_bytes = expected & kChargedBytesMask; + const uint64_t effective_quota_bytes = + effective_quota_bytes_.load(std::memory_order_acquire); + if (bytes != 0 && (charged_bytes > effective_quota_bytes || + bytes > effective_quota_bytes - charged_bytes)) { + const uint64_t sequence_after = + policy_sequence_.load(std::memory_order_acquire); + if (sequence_before == sequence_after && !(sequence_after & 1)) { + const unsigned __int128 demand = + static_cast(charged_bytes) + bytes; + const unsigned __int128 deficit = + demand - effective_quota_bytes; + return ChargeFailure( + TenantQuotaError::kQuotaExceeded, + deficit > std::numeric_limits::max() + ? std::numeric_limits::max() + : static_cast(deficit)); + } + continue; + } + + if (bytes != 0) { + const uint64_t desired = charged_bytes + bytes; + if (!charged_state_.compare_exchange_weak( + expected, desired, std::memory_order_acq_rel, + std::memory_order_acquire)) { + continue; + } + } + + const uint64_t state_after = + charged_state_.load(std::memory_order_acquire); + const uint64_t sequence_after = + policy_sequence_.load(std::memory_order_acquire); + if (sequence_before == sequence_after && !(sequence_after & 1) && + !(state_after & kAdmissionClosed)) { + return {}; + } + + if (bytes != 0) { + auto release_result = Release(bytes); + if (!release_result) { + return ChargeFailure(release_result.error()); + } + } + } } -} // namespace +TenantQuotaResult TenantQuotaAccount::Release(uint64_t bytes) { + if (bytes == 0) { + return {}; + } + + uint64_t expected = charged_state_.load(std::memory_order_acquire); + for (;;) { + const uint64_t charged_bytes = expected & kChargedBytesMask; + if (bytes > charged_bytes) { + return AccountingMismatch(); + } + + const uint64_t desired = + (expected & kAdmissionClosed) | (charged_bytes - bytes); + if (charged_state_.compare_exchange_weak(expected, desired, + std::memory_order_acq_rel, + std::memory_order_acquire)) { + return {}; + } + } +} + +uint64_t TenantQuotaAccount::ChargedBytes() const { + return charged_state_.load(std::memory_order_acquire) & kChargedBytesMask; +} + +uint64_t TenantQuotaAccount::EffectiveQuotaBytes() const { + return effective_quota_bytes_.load(std::memory_order_acquire); +} + +bool TenantQuotaAccount::AdmissionClosed() const { + return charged_state_.load(std::memory_order_acquire) & kAdmissionClosed; +} + +void TenantQuotaAccount::BeginPolicyUpdate() { + const uint64_t previous = + policy_sequence_.fetch_add(1, std::memory_order_acq_rel); + assert((previous & 1) == 0); + (void)previous; +} + +void TenantQuotaAccount::EndPolicyUpdate() { + const uint64_t previous = + policy_sequence_.fetch_add(1, std::memory_order_release); + assert((previous & 1) != 0); + (void)previous; +} + +void TenantQuotaAccount::SetAdmissionClosed(bool closed) { + uint64_t expected = charged_state_.load(std::memory_order_acquire); + for (;;) { + const uint64_t desired = + closed ? expected | kAdmissionClosed : expected & kChargedBytesMask; + if (charged_state_.compare_exchange_weak(expected, desired, + std::memory_order_acq_rel, + std::memory_order_acquire)) { + return; + } + } +} + +void TenantQuotaAccount::ApplyEffectiveQuota(uint64_t effective_quota_bytes) { + effective_quota_bytes_.store(effective_quota_bytes, + std::memory_order_release); + SetAdmissionClosed(!has_explicit_policy_); +} + +void TenantQuotaAccount::SetChargedBytesForRebuild(uint64_t charged_bytes) { + assert(charged_bytes <= kMaxChargedBytes); + charged_state_.store(charged_bytes | kAdmissionClosed, + std::memory_order_release); +} std::map TenantQuotaTable::BuildEffectiveQuotaAssignments( const std::vector& tenants, @@ -93,79 +238,87 @@ std::map TenantQuotaTable::BuildEffectiveQuotaAssignments( TenantQuotaResult TenantQuotaTable::UpsertTenantPolicy( const TenantId& tenant_id, uint64_t requested_quota_bytes) { - if (requested_quota_bytes == 0) { + if (requested_quota_bytes == 0 || + requested_quota_bytes > TenantQuotaAccount::kMaxChargedBytes) { return tl::make_unexpected(TenantQuotaError::kInvalidArgument); } - auto& state = GetOrCreateState(tenant_id); - state.requested_quota_bytes = requested_quota_bytes; - state.has_explicit_policy = true; - RefreshOverQuota(&state); + auto& account = GetOrCreateAccount(tenant_id); + account.BeginPolicyUpdate(); + account.SetAdmissionClosed(true); + account.requested_quota_bytes_ = requested_quota_bytes; + account.has_explicit_policy_ = true; + account.EndPolicyUpdate(); return {}; } -TenantQuotaPolicyResult TenantQuotaTable::DisableTenantPolicyIfEmpty( +TenantQuotaResult TenantQuotaTable::DisableTenantPolicyIfEmpty( const TenantId& tenant_id) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || !it->second.has_explicit_policy) { + auto it = accounts_.find(tenant_id); + if (it == accounts_.end() || !it->second->has_explicit_policy_) { return tl::make_unexpected(TenantQuotaError::kTenantNotFound); } - auto& state = it->second; - if (state.used_bytes != 0 || state.reserved_bytes != 0 || - state.committed_count != 0 || state.metadata_object_count != 0) { + auto& account = *it->second; + account.BeginPolicyUpdate(); + account.SetAdmissionClosed(true); + if (account.ChargedBytes() != 0) { + account.SetAdmissionClosed(false); + account.EndPolicyUpdate(); return tl::make_unexpected(TenantQuotaError::kTenantNotEmpty); } - const uint64_t requested_quota_bytes = state.requested_quota_bytes; - state.requested_quota_bytes = 0; - state.effective_quota_bytes = 0; - state.has_explicit_policy = false; - RefreshOverQuota(&state); - EraseIfLazyEmpty(it); - return requested_quota_bytes; + account.requested_quota_bytes_ = 0; + account.effective_quota_bytes_.store(0, std::memory_order_release); + account.has_explicit_policy_ = false; + account.EndPolicyUpdate(); + return {}; } -void TenantQuotaTable::ApplyTenantPolicies( +TenantQuotaResult TenantQuotaTable::ApplyTenantPolicies( const TenantQuotaPolicyMap& policies) { - for (auto it = tenants_.begin(); it != tenants_.end();) { - auto policy_it = policies.find(it->first); - auto& state = it->second; - if (policy_it != policies.end()) { - state.requested_quota_bytes = policy_it->second; - state.has_explicit_policy = true; - RefreshOverQuota(&state); - ++it; - continue; + for (const auto& [_, requested_quota_bytes] : policies) { + if (requested_quota_bytes == 0 || + requested_quota_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return tl::make_unexpected(TenantQuotaError::kInvalidArgument); } + } - state.requested_quota_bytes = 0; - state.effective_quota_bytes = 0; - state.has_explicit_policy = false; - if (IsLazyEmptyTenant(state)) { - it = tenants_.erase(it); + for (auto& [tenant_id, account_ptr] : accounts_) { + auto& account = *account_ptr; + auto policy_it = policies.find(tenant_id); + account.BeginPolicyUpdate(); + account.SetAdmissionClosed(true); + if (policy_it == policies.end()) { + account.requested_quota_bytes_ = 0; + account.effective_quota_bytes_.store(0, std::memory_order_release); + account.has_explicit_policy_ = false; } else { - RefreshOverQuota(&state); - ++it; + account.requested_quota_bytes_ = policy_it->second; + account.has_explicit_policy_ = true; } + account.EndPolicyUpdate(); } for (const auto& [tenant_id, requested_quota_bytes] : policies) { - auto [it, inserted] = tenants_.try_emplace(tenant_id); - if (!inserted) { + auto& account = GetOrCreateAccount(tenant_id); + if (account.has_explicit_policy_) { continue; } - it->second.requested_quota_bytes = requested_quota_bytes; - it->second.has_explicit_policy = true; - RefreshOverQuota(&it->second); + account.BeginPolicyUpdate(); + account.SetAdmissionClosed(true); + account.requested_quota_bytes_ = requested_quota_bytes; + account.has_explicit_policy_ = true; + account.EndPolicyUpdate(); } + return {}; } TenantQuotaPolicyMap TenantQuotaTable::GetTenantPolicies() const { TenantQuotaPolicyMap policies; - for (const auto& [tenant_id, state] : tenants_) { - if (state.has_explicit_policy) { - policies.emplace(tenant_id, state.requested_quota_bytes); + for (const auto& [tenant_id, account] : accounts_) { + if (account->has_explicit_policy_) { + policies.emplace(tenant_id, account->requested_quota_bytes_); } } return policies; @@ -178,259 +331,97 @@ void TenantQuotaTable::RecomputeEffectiveQuotas( } bool TenantQuotaTable::IsTenantRegistered(const TenantId& tenant_id) const { - auto it = tenants_.find(tenant_id); - return it != tenants_.end() && it->second.has_explicit_policy; + auto it = accounts_.find(tenant_id); + return it != accounts_.end() && it->second->has_explicit_policy_; +} + +TenantQuotaHandle TenantQuotaTable::GetOrCreateTenantHandle( + const TenantId& tenant_id) { + return &GetOrCreateAccount(tenant_id); } std::optional TenantQuotaTable::GetTenantSnapshot( const TenantId& tenant_id) const { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end()) { + auto it = accounts_.find(tenant_id); + if (it == accounts_.end() || (!it->second->has_explicit_policy_ && + it->second->ChargedBytes() == 0)) { return std::nullopt; } - return MakeSnapshot(it->first, it->second); + return MakeSnapshot(it->first, *it->second); } std::vector TenantQuotaTable::ListTenantSnapshots() const { std::vector snapshots; - snapshots.reserve(tenants_.size()); - for (const auto& [tenant_id, state] : tenants_) { - if (!IsLazyEmptyTenant(state)) { - snapshots.push_back(MakeSnapshot(tenant_id, state)); + snapshots.reserve(accounts_.size()); + for (const auto& [tenant_id, account] : accounts_) { + if (!account->has_explicit_policy_ && account->ChargedBytes() == 0) { + continue; } + snapshots.push_back(MakeSnapshot(tenant_id, *account)); } return snapshots; } -uint64_t TenantQuotaTable::ComputeDeficit(const TenantId& tenant_id, - uint64_t incoming_bytes) const { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end()) { - return incoming_bytes; - } - - const auto& state = it->second; - const unsigned __int128 demand = - static_cast(state.used_bytes) + - state.reserved_bytes + incoming_bytes; - if (demand <= state.effective_quota_bytes) { - return 0; - } - - const unsigned __int128 deficit = demand - state.effective_quota_bytes; - return deficit > std::numeric_limits::max() - ? std::numeric_limits::max() - : static_cast(deficit); -} - -TenantQuotaResult TenantQuotaTable::Reserve(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || !it->second.has_explicit_policy) { - return tl::make_unexpected(TenantQuotaError::kTenantNotRegistered); - } - if (bytes == 0) { - return {}; - } - - auto& state = it->second; - if (static_cast(state.used_bytes) + - state.reserved_bytes + bytes > - state.effective_quota_bytes) { - return tl::make_unexpected(TenantQuotaError::kQuotaExceeded); - } - - state.reserved_bytes += bytes; - RefreshOverQuota(&state); - return {}; -} - -TenantQuotaResult TenantQuotaTable::Commit(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.reserved_bytes < bytes) { - return AccountingMismatch(); - } - if (bytes == 0) { - return {}; - } - - auto& state = it->second; - state.reserved_bytes -= bytes; - state.used_bytes = SaturatingAdd(state.used_bytes, bytes); - if (state.committed_count < std::numeric_limits::max()) { - ++state.committed_count; - } - RefreshOverQuota(&state); - return {}; -} - -TenantQuotaResult TenantQuotaTable::CommitAdditional(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.reserved_bytes < bytes) { - return AccountingMismatch(); - } - if (bytes == 0) { - return {}; +TenantQuotaResult TenantQuotaTable::RebuildUsage( + const TenantQuotaUsageMap& usage) { + for (const auto& [_, charged_bytes] : usage) { + if (charged_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return tl::make_unexpected(TenantQuotaError::kInvalidArgument); + } } - auto& state = it->second; - state.reserved_bytes -= bytes; - state.used_bytes = SaturatingAdd(state.used_bytes, bytes); - RefreshOverQuota(&state); - return {}; -} - -TenantQuotaResult TenantQuotaTable::Abort(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.reserved_bytes < bytes) { - return AccountingMismatch(); + for (auto& [_, account] : accounts_) { + account->BeginPolicyUpdate(); + account->SetChargedBytesForRebuild(0); + account->EndPolicyUpdate(); } - if (bytes == 0) { - return {}; - } - - it->second.reserved_bytes -= bytes; - RefreshOverQuota(&it->second); - EraseIfLazyEmpty(it); - return {}; -} -TenantQuotaResult TenantQuotaTable::Release(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.used_bytes < bytes || - (bytes != 0 && it->second.committed_count == 0)) { - return AccountingMismatch(); - } - if (bytes == 0) { - return {}; - } - - auto& state = it->second; - state.used_bytes -= bytes; - --state.committed_count; - RefreshOverQuota(&state); - EraseIfLazyEmpty(it); - return {}; -} - -TenantQuotaResult TenantQuotaTable::ReleasePartial(const TenantId& tenant_id, - uint64_t bytes) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.used_bytes < bytes) { - return AccountingMismatch(); - } - if (bytes == 0) { - return {}; + for (const auto& [tenant_id, charged_bytes] : usage) { + auto& account = GetOrCreateAccount(tenant_id); + if (!account.has_explicit_policy_) { + account.requested_quota_bytes_ = 0; + account.effective_quota_bytes_.store(0, std::memory_order_release); + } + account.BeginPolicyUpdate(); + account.SetChargedBytesForRebuild(charged_bytes); + account.EndPolicyUpdate(); } - - it->second.used_bytes -= bytes; - RefreshOverQuota(&it->second); - EraseIfLazyEmpty(it); return {}; } -void TenantQuotaTable::IncrementMetadataObjectCount(const TenantId& tenant_id) { - auto& state = GetOrCreateState(tenant_id); - if (state.metadata_object_count < std::numeric_limits::max()) { - ++state.metadata_object_count; - } - RefreshOverQuota(&state); -} - -TenantQuotaResult TenantQuotaTable::DecrementMetadataObjectCount( +TenantQuotaAccount& TenantQuotaTable::GetOrCreateAccount( const TenantId& tenant_id) { - auto it = tenants_.find(tenant_id); - if (it == tenants_.end() || it->second.metadata_object_count == 0) { - return AccountingMismatch(); - } - - --it->second.metadata_object_count; - RefreshOverQuota(&it->second); - EraseIfLazyEmpty(it); - return {}; -} - -void TenantQuotaTable::RebuildUsage(const TenantQuotaUsageMap& usage) { - for (auto& [_, state] : tenants_) { - state.used_bytes = 0; - state.reserved_bytes = 0; - state.committed_count = 0; - state.metadata_object_count = 0; - } - - for (const auto& [tenant_id, tenant_usage] : usage) { - auto& state = GetOrCreateState(tenant_id); - if (!state.has_explicit_policy) { - state.requested_quota_bytes = 0; - state.effective_quota_bytes = 0; - } - state.used_bytes = tenant_usage.used_bytes; - state.committed_count = tenant_usage.committed_count; - state.metadata_object_count = tenant_usage.metadata_object_count; - RefreshOverQuota(&state); - } - - for (auto it = tenants_.begin(); it != tenants_.end();) { - if (IsLazyEmptyTenant(it->second)) { - it = tenants_.erase(it); - } else { - RefreshOverQuota(&it->second); - ++it; - } + auto [it, inserted] = accounts_.try_emplace(tenant_id); + if (inserted) { + it->second = std::make_unique(); } -} - -TenantQuotaTable::TenantQuotaState& TenantQuotaTable::GetOrCreateState( - const TenantId& tenant_id) { - return tenants_.try_emplace(tenant_id).first->second; + return *it->second; } TenantQuotaSnapshot TenantQuotaTable::MakeSnapshot( - const TenantId& tenant_id, const TenantQuotaState& state) const { + const TenantId& tenant_id, const TenantQuotaAccount& account) const { + const uint64_t charged_bytes = account.ChargedBytes(); + const uint64_t effective_quota_bytes = account.EffectiveQuotaBytes(); return TenantQuotaSnapshot{ .tenant_id = tenant_id, - .requested_quota_bytes = state.requested_quota_bytes, - .effective_quota_bytes = state.effective_quota_bytes, - .used_bytes = state.used_bytes, - .reserved_bytes = state.reserved_bytes, - .committed_count = state.committed_count, - .metadata_object_count = state.metadata_object_count, - .has_explicit_policy = state.has_explicit_policy, - .over_quota = state.over_quota, + .requested_quota_bytes = account.requested_quota_bytes_, + .effective_quota_bytes = effective_quota_bytes, + .charged_bytes = charged_bytes, + .admission_closed = account.AdmissionClosed(), + .has_explicit_policy = account.has_explicit_policy_, + .over_quota = charged_bytes > effective_quota_bytes, }; } -bool TenantQuotaTable::IsLazyEmptyTenant(const TenantQuotaState& state) { - return !state.has_explicit_policy && state.used_bytes == 0 && - state.reserved_bytes == 0 && state.committed_count == 0 && - state.metadata_object_count == 0; -} - -void TenantQuotaTable::RefreshOverQuota(TenantQuotaState* state) { - state->over_quota = - (!state->has_explicit_policy && state->metadata_object_count > 0) || - static_cast(state->used_bytes) + - state->reserved_bytes > - state->effective_quota_bytes; -} - void TenantQuotaTable::ApplyEffectiveQuotas( const std::map& effective_quotas) { - for (auto& [tenant_id, state] : tenants_) { + for (auto& [tenant_id, account] : accounts_) { auto it = effective_quotas.find(tenant_id); - state.effective_quota_bytes = + const uint64_t effective_quota_bytes = it == effective_quotas.end() ? 0 : it->second; - RefreshOverQuota(&state); - } -} - -void TenantQuotaTable::EraseIfLazyEmpty(StateMap::iterator it) { - if (IsLazyEmptyTenant(it->second)) { - tenants_.erase(it); + account->BeginPolicyUpdate(); + account->ApplyEffectiveQuota(effective_quota_bytes); + account->EndPolicyUpdate(); } } diff --git a/mooncake-store/src/tenant_quota_ledger.cpp b/mooncake-store/src/tenant_quota_ledger.cpp new file mode 100644 index 0000000000..33d3dbe7d8 --- /dev/null +++ b/mooncake-store/src/tenant_quota_ledger.cpp @@ -0,0 +1,196 @@ +#include "tenant_quota_ledger.h" + +#include + +namespace mooncake { +namespace { + +TenantQuotaResult InvalidArgument() { + return tl::make_unexpected(TenantQuotaError::kInvalidArgument); +} + +TenantQuotaResult AccountingMismatch() { + return tl::make_unexpected(TenantQuotaError::kAccountingMismatch); +} + +bool AddOverflows(uint64_t lhs, uint64_t rhs) { + return rhs > TenantQuotaAccount::kMaxChargedBytes || + lhs > TenantQuotaAccount::kMaxChargedBytes - rhs; +} + +} // namespace + +TenantQuotaResult TenantQuotaLedger::AdoptPendingCharge( + TenantQuotaHandle account, uint64_t bytes) { + if (account == nullptr || bytes > TenantQuotaAccount::kMaxChargedBytes) { + return InvalidArgument(); + } + if (pending_bytes_ != 0 || AddOverflows(TotalChargedBytes(), bytes)) { + return AccountingMismatch(); + } + pending_bytes_ = bytes; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::SettlePrimaryWrite( + TenantQuotaHandle account, uint64_t actual_committed_bytes) { + if (account == nullptr || + actual_committed_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return InvalidArgument(); + } + if (AddOverflows(pending_bytes_, committed_bytes_)) { + return AccountingMismatch(); + } + const uint64_t primary_bytes = pending_bytes_ + committed_bytes_; + if (actual_committed_bytes > primary_bytes || + AddOverflows(primary_bytes, replaced_bytes_)) { + return AccountingMismatch(); + } + + const uint64_t total_bytes = primary_bytes + replaced_bytes_; + auto release_result = + account->Release(total_bytes - actual_committed_bytes); + if (!release_result) { + return release_result; + } + pending_bytes_ = 0; + committed_bytes_ = actual_committed_bytes; + replaced_bytes_ = 0; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::SettleAdditional(TenantQuotaHandle account, + uint64_t pending_bytes, + uint64_t actual_bytes) { + if (account == nullptr || + pending_bytes > TenantQuotaAccount::kMaxChargedBytes || + actual_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return InvalidArgument(); + } + if (actual_bytes > pending_bytes || + AddOverflows(committed_bytes_, actual_bytes) || + AddOverflows(TotalChargedBytes(), actual_bytes)) { + return AccountingMismatch(); + } + + auto release_result = account->Release(pending_bytes - actual_bytes); + if (!release_result) { + return release_result; + } + committed_bytes_ += actual_bytes; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::RefundPending(TenantQuotaHandle account) { + if (account == nullptr) { + return InvalidArgument(); + } + if (pending_bytes_ == 0) { + return AccountingMismatch(); + } + auto release_result = account->Release(pending_bytes_); + if (!release_result) { + return release_result; + } + pending_bytes_ = 0; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::ReleaseCommitted(TenantQuotaHandle account, + uint64_t bytes) { + if (account == nullptr || bytes > TenantQuotaAccount::kMaxChargedBytes) { + return InvalidArgument(); + } + if (bytes == 0) { + return {}; + } + if (bytes > committed_bytes_) { + return AccountingMismatch(); + } + auto release_result = account->Release(bytes); + if (!release_result) { + return release_result; + } + committed_bytes_ -= bytes; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::TransferReplacementCharge( + TenantQuotaHandle account, TenantQuotaLedger& destination) { + if (account == nullptr || this == &destination) { + return InvalidArgument(); + } + if (pending_bytes_ != 0 || destination.replaced_bytes_ != 0 || + AddOverflows(committed_bytes_, replaced_bytes_)) { + return AccountingMismatch(); + } + + const uint64_t transfer_bytes = committed_bytes_ + replaced_bytes_; + if (transfer_bytes == 0 || + AddOverflows(destination.TotalChargedBytes(), transfer_bytes)) { + return AccountingMismatch(); + } + destination.replaced_bytes_ = transfer_bytes; + committed_bytes_ = 0; + replaced_bytes_ = 0; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::ReleaseReplacement( + TenantQuotaHandle account) { + if (account == nullptr) { + return InvalidArgument(); + } + if (replaced_bytes_ == 0) { + return AccountingMismatch(); + } + auto release_result = account->Release(replaced_bytes_); + if (!release_result) { + return release_result; + } + replaced_bytes_ = 0; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::ReleaseAll(TenantQuotaHandle account) { + if (account == nullptr) { + return InvalidArgument(); + } + const uint64_t total_bytes = TotalChargedBytes(); + if (total_bytes == 0) { + return AccountingMismatch(); + } + auto release_result = account->Release(total_bytes); + if (!release_result) { + return release_result; + } + pending_bytes_ = 0; + committed_bytes_ = 0; + replaced_bytes_ = 0; + return {}; +} + +TenantQuotaResult TenantQuotaLedger::Rebuild(TenantQuotaHandle account, + uint64_t committed_bytes) { + if (account == nullptr || + committed_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return InvalidArgument(); + } + pending_bytes_ = 0; + committed_bytes_ = committed_bytes; + replaced_bytes_ = 0; + return {}; +} + +uint64_t TenantQuotaLedger::TotalChargedBytes() const { + if (AddOverflows(pending_bytes_, committed_bytes_)) { + return std::numeric_limits::max(); + } + const uint64_t pending_and_committed = pending_bytes_ + committed_bytes_; + if (AddOverflows(pending_and_committed, replaced_bytes_)) { + return std::numeric_limits::max(); + } + return pending_and_committed + replaced_bytes_; +} + +} // namespace mooncake diff --git a/mooncake-store/src/tenant_quota_policy_store.cpp b/mooncake-store/src/tenant_quota_policy_store.cpp index 1b3db9147a..5420f15020 100644 --- a/mooncake-store/src/tenant_quota_policy_store.cpp +++ b/mooncake-store/src/tenant_quota_policy_store.cpp @@ -1,5 +1,7 @@ #include "tenant_quota_policy_store.h" +#include "tenant_quota.h" + #include #include #include @@ -211,7 +213,12 @@ tl::expected ParseTenantQuotaBytes( if (number > std::numeric_limits::max() / multiplier) { return tl::make_unexpected("quota byte value overflows uint64"); } - return number * multiplier; + const uint64_t quota_bytes = number * multiplier; + if (quota_bytes > TenantQuotaAccount::kMaxChargedBytes) { + return tl::make_unexpected( + "quota exceeds maximum supported tenant charge"); + } + return quota_bytes; } tl::expected ParseTenantQuotaPolicyYaml( diff --git a/mooncake-store/tests/CMakeLists.txt b/mooncake-store/tests/CMakeLists.txt index a5855d69d5..7e951d0d37 100644 --- a/mooncake-store/tests/CMakeLists.txt +++ b/mooncake-store/tests/CMakeLists.txt @@ -127,6 +127,7 @@ if(USE_CUDA) endif() endif() add_store_test(tenant_quota_test tenant_quota_test.cpp) +add_store_test(tenant_quota_ledger_test tenant_quota_ledger_test.cpp) add_store_test(tenant_id_test tenant_id_test.cpp) add_store_test(segment_test segment_test.cpp) add_store_test(offset_allocator_test offset_allocator_test.cpp) diff --git a/mooncake-store/tests/ha/master_service_ha_test.cpp b/mooncake-store/tests/ha/master_service_ha_test.cpp index 7e733a38b1..6cb0b00ef6 100644 --- a/mooncake-store/tests/ha/master_service_ha_test.cpp +++ b/mooncake-store/tests/ha/master_service_ha_test.cpp @@ -406,7 +406,7 @@ class MasterServiceHATest : public ::testing::Test { static uint64_t TenantUsedBytes(MasterService& service) { auto snapshot = service.GetTenantQuotaSnapshot(kDefaultTenant); EXPECT_TRUE(snapshot.has_value()); - return snapshot ? snapshot->used_bytes : 0; + return snapshot ? snapshot->charged_bytes : 0; } void ReadRemoveBatchEventually(OpLogBatchStorage& storage, @@ -519,7 +519,8 @@ class MasterServiceHATest : public ::testing::Test { const size_t shard_idx = service->getMetadataShardIndex(tenant, key); auto shard_access = MasterService::MetadataShardAccessorRW(service, shard_idx); - auto& tenant_state = shard_access->tenants[tenant]; + auto& tenant_state = + service->GetOrCreateTenantState(shard_access.get(), tenant); tenant_state.promotion_tasks.emplace( key, MasterService::PromotionTask{ .source_id = 0, @@ -718,6 +719,32 @@ TEST_F(MasterServiceHATest, RestoreFromStandbyPreservesMemoryBufferDescriptor) { EXPECT_EQ(restored.transport_endpoint_, endpoint); } +TEST_F(MasterServiceHATest, RestoreFromStandbyRebuildsTenantQuotaAccounting) { + const TenantId tenant_id("tenant_a"); + constexpr uint64_t object_size = 1024; + auto config = MasterServiceConfig::builder() + .set_enable_multi_tenants(true) + .set_tenant_quota_connector_type("file") + .set_tenant_quota_connector_uri(WriteTenantPolicyFile( + {{tenant_id.value(), object_size}})) + .build(); + MasterService service(config); + + const std::string key = "standby_quota_key"; + const std::string endpoint = "standby_quota_segment"; + auto object = MakeStandbyObject(key, endpoint, object_size); + object.tenant_id = tenant_id.value(); + + service.RestoreFromStandbySnapshot({object}, 7, + {MakeStandbyMemorySegment(endpoint)}); + + auto snapshot = service.GetTenantQuotaSnapshot(tenant_id); + ASSERT_TRUE(snapshot.has_value()); + EXPECT_EQ(snapshot->charged_bytes, object_size); + ASSERT_TRUE(service.Remove(key, tenant_id, /*force=*/true).has_value()); + EXPECT_EQ(service.GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 0); +} + TEST_F(MasterServiceHATest, RemountMakesRestoredMemoryReplicaReady) { MasterService service( MasterServiceConfig::builder().set_enable_ha(false).build()); @@ -2525,6 +2552,88 @@ TEST_F(MasterServiceHATest, EXPECT_TRUE(after_finalize.has_value()) << toString(after_finalize.error()); } +TEST_F(MasterServiceHATest, + MoveFinalizeDuringSameSizeUpsertReleasesOnlyRemovedReplicaQuota) { + const std::string cluster_id = "test_batch_move_upsert_quota"; + constexpr uint64_t object_size = 1024; + auto backend = std::make_shared(); + auto service_config = + MasterServiceConfig::builder() + .set_default_kv_lease_ttl(50) + .set_enable_ha(true) + .set_enable_oplog(true) + .set_cluster_id(cluster_id) + .set_oplog_batch_max_entries(1) + .set_eviction_high_watermark_ratio(1.0) + .set_enable_multi_tenants(true) + .set_tenant_quota_connector_type("file") + .set_tenant_quota_connector_uri(WriteTenantPolicyFile( + {{kDefaultTenant.value(), 2 * object_size}})) + .build(); + MasterService service(service_config); + ASSERT_EQ(ErrorCode::OK, service.SetBatchOpLogBackendForTesting(backend)); + + const std::string source_name = "batch_move_upsert_quota_src"; + const std::string target_name = "batch_move_upsert_quota_dst"; + auto source = PrepareSimpleSegment(service, source_name, + kDefaultSegmentBase, object_size); + OpLogBatchStorage storage(cluster_id, *backend); + OpLogBatchRecord batch; + ReadBatchEventually(storage, 1, batch); + PrepareSimpleSegment(service, target_name, + kDefaultSegmentBase + kDefaultSegmentSize, + object_size); + ReadBatchEventually(storage, 2, batch); + + const std::string key = "batch_move_upsert_quota_key"; + PutObjectOnSegment(service, source.client_id, key, source_name, + object_size); + ReadBatchEventually(storage, 3, batch); + ASSERT_TRUE(service + .MoveStart(source.client_id, key, kDefaultTenant, + source_name, target_name) + .has_value()); + + backend->BlockTxn(); + auto move_future = std::async(std::launch::async, [&] { + return service.MoveEnd(source.client_id, key, kDefaultTenant); + }); + const auto status = move_future.wait_for(std::chrono::milliseconds(200)); + EXPECT_EQ(std::future_status::ready, status); + if (status != std::future_status::ready) { + backend->AllowTxn(); + } + auto move_end = move_future.get(); + ASSERT_TRUE(move_end.has_value()) << toString(move_end.error()); + ASSERT_EQ(service.GetTenantQuotaSnapshot(kDefaultTenant)->charged_bytes, + 2 * object_size); + + ReplicateConfig config; + config.replica_num = 1; + config.preferred_segments = {target_name}; + auto upsert = service.UpsertStart(source.client_id, key, kDefaultTenant, + object_size, config); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + + backend->AllowTxn(); + ReadBatchEventually(storage, 4, batch); + uint64_t charged_bytes = 2 * object_size; + for (int i = 0; i < 50 && charged_bytes == 2 * object_size; ++i) { + charged_bytes = + service.GetTenantQuotaSnapshot(kDefaultTenant)->charged_bytes; + if (charged_bytes == 2 * object_size) { + std::this_thread::sleep_for(std::chrono::milliseconds(20)); + } + } + ASSERT_EQ(charged_bytes, object_size); + + auto upsert_end = service.UpsertEnd(source.client_id, key, kDefaultTenant, + ReplicaType::MEMORY); + ASSERT_TRUE(upsert_end.has_value()) << toString(upsert_end.error()); + EXPECT_EQ(service.GetTenantQuotaSnapshot(kDefaultTenant)->charged_bytes, + object_size); +} + TEST_F(MasterServiceHATest, NotifyOffloadSuccessFallbackWritesBatchRecordOpLog) { const std::string cluster_id = "test_batch_record_offload_cluster"; diff --git a/mooncake-store/tests/master_admin_server_test.cpp b/mooncake-store/tests/master_admin_server_test.cpp index 8c572a4199..3b2b9b2b2c 100644 --- a/mooncake-store/tests/master_admin_server_test.cpp +++ b/mooncake-store/tests/master_admin_server_test.cpp @@ -535,7 +535,8 @@ TEST_F(MasterAdminServerTest, TenantQuotaAdminLifecycleEndpoints) { auto one = HttpGet(port, "/api/v1/tenant_quotas?tenant_id=tenant-a"); EXPECT_EQ(one.http_status, 200); - EXPECT_NE(one.body.find("\"committed_count\":0"), std::string::npos); + EXPECT_NE(one.body.find("\"charged_bytes\":0"), std::string::npos); + EXPECT_NE(one.body.find("\"admission_closed\":false"), std::string::npos); EXPECT_NE(one.body.find("\"over_quota\":false"), std::string::npos); ReplicateConfig cfg; @@ -549,6 +550,21 @@ TEST_F(MasterAdminServerTest, TenantQuotaAdminLifecycleEndpoints) { ReplicaType::MEMORY, "tenant-a") .has_value()); + auto metrics = HttpGet(port, "/metrics"); + EXPECT_EQ(metrics.http_status, 200); + EXPECT_NE( + metrics.body.find( + "mooncake_tenant_quota_charged_bytes{tenant_id=\"tenant-a\"} 100"), + std::string::npos); + EXPECT_NE(metrics.body.find( + "mooncake_tenant_quota_admission_closed{tenant_id=\"tenant-a" + "\"} 0"), + std::string::npos); + EXPECT_EQ(metrics.body.find("mooncake_tenant_quota_reserved_bytes"), + std::string::npos); + EXPECT_EQ(metrics.body.find("mooncake_tenant_quota_used_bytes"), + std::string::npos); + auto delete_non_empty = HttpDelete(port, "/api/v1/tenant_quotas?tenant_id=tenant-a"); EXPECT_EQ(delete_non_empty.http_status, 409); @@ -598,6 +614,11 @@ TEST_F(MasterAdminServerTest, TenantQuotaAdminValidationErrors) { "{\"requested_quota_bytes\":0}"); EXPECT_EQ(zero_explicit.http_status, 400); + auto above_atomic_range = + HttpPutJson(port, "/api/v1/tenant_quotas?tenant_id=tenant-a", + "{\"requested_quota_bytes\":9223372036854775808}"); + EXPECT_EQ(above_atomic_range.http_status, 400); + auto reserved_tenant = HttpGet(port, "/api/v1/tenant_quotas?tenant_id=_system"); EXPECT_EQ(reserved_tenant.http_status, 400); diff --git a/mooncake-store/tests/master_service_tenant_quota_test.cpp b/mooncake-store/tests/master_service_tenant_quota_test.cpp index 5016af2b55..2575371ec0 100644 --- a/mooncake-store/tests/master_service_tenant_quota_test.cpp +++ b/mooncake-store/tests/master_service_tenant_quota_test.cpp @@ -233,9 +233,94 @@ class MasterServiceTenantQuotaTest : public ::testing::Test { } #endif - tl::expected ReserveTenantQuotaForTest( + tl::expected ChargeTenantQuotaForTest( MasterService& service, const TenantId& tenant_id, uint64_t bytes) { - return service.ReserveTenantQuota(tenant_id, bytes); + return service.ChargeTenantQuota( + service.tenant_quota_table_.GetOrCreateTenantHandle(tenant_id), + bytes); + } + + TenantQuotaHandle GetOrCreateTenantStateHandleForTest( + MasterService& service, size_t shard_idx, const TenantId& tenant_id) { + MasterService::MetadataShardAccessorRW shard(&service, shard_idx); + auto& tenant_state = + service.GetOrCreateTenantState(shard.get(), tenant_id); + return service.GetBoundTenantQuotaHandle(tenant_state); + } + + tl::expected ChargeBoundTenantQuotaForTest( + MasterService& service, TenantQuotaHandle account, uint64_t bytes) { + return service.ChargeTenantQuota(account, bytes); + } + + void ReleaseBoundTenantQuotaForTest(MasterService& service, + TenantQuotaHandle account, + uint64_t bytes) { + service.ReleaseTenantQuota(account, bytes); + } + + void DiscardExpiredProcessingForTest(MasterService& service, + const TenantId& tenant_id, + const std::string& key) { + const size_t shard_idx = service.getMetadataShardIndex(tenant_id, key); + MasterService::MetadataShardAccessorRW shard(&service, shard_idx); + service.DiscardExpiredProcessingReplicas( + shard, std::chrono::system_clock::time_point::max()); + } + + void FinalizeExpiredProcessingForTest(MasterService& service, + const TenantId& tenant_id, + const std::string& key) { + OpLogEntry entry; + entry.tenant_id = tenant_id.value(); + entry.object_key = key; + service.FinalizeExpiredProcessingReplicasAfterDurable( + entry, std::chrono::system_clock::now()); + } + + void FinalizeRemovedMemoryReplicasForTest(MasterService& service, + const TenantId& tenant_id, + const std::string& key) { + std::vector removed_ids; + { + MasterService::MetadataAccessorRW accessor( + &service, MasterService::ObjectIdentity{tenant_id, key}); + ASSERT_TRUE(accessor.Exists()); + accessor.Get().VisitReplicas( + &Replica::fn_is_memory_replica, + [&removed_ids](Replica& replica) { + removed_ids.push_back(replica.id()); + replica.mark_removed(); + }); + } + ASSERT_FALSE(removed_ids.empty()); + + OpLogEntry entry; + entry.tenant_id = tenant_id.value(); + entry.object_key = key; + service.FinalizeRemovedReplicasAfterDurable( + entry, removed_ids, MasterService::QuotaEraseMode::kFull); + } + + void AddCompletedDiskReplica(MasterService& service, const UUID& client_id, + const std::string& key, + const TenantId& tenant_id, uint64_t size) { + Replica disk_replica(client_id, size, "disk-endpoint", + ReplicaStatus::COMPLETE); + auto result = + service.AddReplica(client_id, key, tenant_id, disk_replica); + ASSERT_TRUE(result.has_value()) << toString(result.error()); + } + + void ExpectDiskOnlyObjectAndChargedBytes(MasterService& service, + const TenantId& tenant_id, + const std::string& key, + uint64_t charged_bytes) { + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, charged_bytes); + auto replicas = service.GetReplicaList(key, tenant_id); + ASSERT_TRUE(replicas.has_value()) << toString(replicas.error()); + ASSERT_EQ(replicas->replicas.size(), 1); + EXPECT_TRUE(replicas->replicas.front().is_local_disk_replica()); } std::unique_lock LockSnapshotForTest( @@ -249,6 +334,11 @@ class MasterServiceTenantQuotaTest : public ::testing::Test { service.tenant_quota_recompute_mutex_); } + std::unique_lock LockTenantQuotaPolicyForTest( + MasterService& service) { + return std::unique_lock(service.tenant_quota_policy_mutex_); + } + ErrorCode MountSegmentWithoutQuotaRecomputeForTest(MasterService& service, size_t size, std::string name) { @@ -339,6 +429,40 @@ TEST_F(MasterServiceTenantQuotaTest, PutComplete(service, client_id, "ok", TenantId("tenant-a"), 10); } +TEST_F(MasterServiceTenantQuotaTest, + SameTenantStatesAcrossMetadataShardsShareBoundHandle) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + MountSegment(service); + + auto* first_handle = + GetOrCreateTenantStateHandleForTest(service, 0, tenant_id); + auto* second_handle = + GetOrCreateTenantStateHandleForTest(service, 1, tenant_id); + + ASSERT_NE(first_handle, nullptr); + EXPECT_EQ(first_handle, second_handle); + + auto charge = ChargeBoundTenantQuotaForTest(service, first_handle, 128); + ASSERT_TRUE(charge.has_value()) << toString(charge.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 128); + + ReleaseBoundTenantQuotaForTest(service, second_handle, 128); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + ChargeRejectsMissingHandleWhenQuotaIsEnabled) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + MountSegment(service); + + auto charge = ChargeBoundTenantQuotaForTest(service, nullptr, 1); + ASSERT_FALSE(charge.has_value()); + EXPECT_EQ(charge.error(), ErrorCode::INTERNAL_ERROR); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 0); +} + TEST_F(MasterServiceTenantQuotaTest, MultiTenantModeRejectsUnregisteredOffloadSuccess) { MasterService service(MakeConfig({{TenantId("tenant-a"), 1000}})); @@ -376,11 +500,11 @@ TEST_F(MasterServiceTenantQuotaTest, auto exists = service.ExistKey("cold", TenantId("tenant-a")); ASSERT_TRUE(exists.has_value()) << toString(exists.error()); EXPECT_TRUE(exists.value()); - EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).used_bytes, 0); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 0); } TEST_F(MasterServiceTenantQuotaTest, - ConnectorPolicyReloadKeepsLocalDiskOnlyOrphanVisible) { + ConnectorPolicyReloadKeepsLocalDiskOnlyOrphanAccessible) { const std::string initial_policy = WritePolicyFile( {{TenantId("tenant-a"), 1000}, {TenantId("tenant-b"), 1000}}); auto config = MasterServiceConfig::builder() @@ -407,12 +531,8 @@ TEST_F(MasterServiceTenantQuotaTest, } ReloadTenantQuotaPolicyFromStore(service); - auto orphan = Snapshot(service, TenantId("tenant-b")); - EXPECT_FALSE(orphan.has_explicit_policy); - EXPECT_EQ(orphan.used_bytes, 0); - EXPECT_EQ(orphan.committed_count, 0); - EXPECT_EQ(orphan.metadata_object_count, 1); - EXPECT_TRUE(orphan.over_quota); + EXPECT_FALSE( + service.GetTenantQuotaSnapshot(TenantId("tenant-b")).has_value()); EXPECT_TRUE(service.Remove("cold", TenantId("tenant-b"), /*force=*/true) .has_value()); @@ -442,7 +562,10 @@ TEST_F(MasterServiceTenantQuotaTest, out << FormatTenantQuotaPolicyYaml(replacement); } ReloadTenantQuotaPolicyFromStore(service); - EXPECT_FALSE(Snapshot(service, TenantId("tenant-b")).has_explicit_policy); + auto orphan = Snapshot(service, TenantId("tenant-b")); + EXPECT_FALSE(orphan.has_explicit_policy); + EXPECT_TRUE(orphan.admission_closed); + EXPECT_EQ(orphan.charged_bytes, 128); StorageObjectMetadata metadata; metadata.data_size = 128; @@ -542,6 +665,7 @@ TEST_F(MasterServiceTenantQuotaTest, auto first = service.PutStart(client_id, "key-a", TenantId("tenant-a"), 80, hard_pinned); ASSERT_TRUE(first.has_value()) << toString(first.error()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 80); ASSERT_TRUE(service .PutEnd(client_id, "key-a", TenantId("tenant-a"), ReplicaType::MEMORY) @@ -552,11 +676,166 @@ TEST_F(MasterServiceTenantQuotaTest, ASSERT_FALSE(over.has_value()); EXPECT_EQ(over.error(), ErrorCode::TENANT_QUOTA_EXCEEDED); - EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).used_bytes, 80); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 80); EXPECT_FALSE( service.GetTenantQuotaSnapshot(TenantId("tenant-b")).has_value()); } +TEST_F(MasterServiceTenantQuotaTest, PutRevokeRefundsStartCharge) { + MasterService service(MakeConfig({{TenantId("tenant-a"), 100}})); + UUID client_id = MountSegment(service); + + auto start = service.PutStart(client_id, "key", TenantId("tenant-a"), 100, + MemoryConfig()); + ASSERT_TRUE(start.has_value()) << toString(start.error()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 100); + + auto over = service.PutStart(client_id, "other", TenantId("tenant-a"), 1, + MemoryConfig()); + ASSERT_FALSE(over.has_value()); + EXPECT_EQ(over.error(), ErrorCode::TENANT_QUOTA_EXCEEDED); + + ASSERT_TRUE(service + .PutRevoke(client_id, "key", TenantId("tenant-a"), + ReplicaType::MEMORY) + .has_value()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + SizeChangingUpsertTransfersAndReleasesReplacementCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + PutComplete(service, client_id, "key", tenant_id, 100); + + auto upsert = + service.UpsertStart(client_id, "key", tenant_id, 200, MemoryConfig()); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 300); + + auto end = + service.UpsertEnd(client_id, "key", tenant_id, ReplicaType::MEMORY); + ASSERT_TRUE(end.has_value()) << toString(end.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 200); +} + +TEST_F(MasterServiceTenantQuotaTest, + SizeChangingUpsertFromDiskOnlyObjectChargesNewReplica) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + + StorageObjectMetadata metadata; + metadata.data_size = 100; + metadata.transport_endpoint = "disk-endpoint"; + std::vector tasks{OffloadTaskItem{ + .tenant_id = tenant_id.value(), .key = "key", .size = 100}}; + ASSERT_TRUE( + service.NotifyOffloadSuccess(client_id, tasks, {metadata}).has_value()); + ASSERT_EQ(Snapshot(service, tenant_id).charged_bytes, 0); + + auto upsert = + service.UpsertStart(client_id, "key", tenant_id, 200, MemoryConfig()); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 200); + + auto end = + service.UpsertEnd(client_id, "key", tenant_id, ReplicaType::MEMORY); + ASSERT_TRUE(end.has_value()) << toString(end.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 200); +} + +TEST_F(MasterServiceTenantQuotaTest, + SizeChangingUpsertRevokeReleasesNewAndReplacementCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + PutComplete(service, client_id, "key", tenant_id, 100); + + auto upsert = + service.UpsertStart(client_id, "key", tenant_id, 200, MemoryConfig()); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 300); + + auto revoke = + service.UpsertRevoke(client_id, "key", tenant_id, ReplicaType::MEMORY); + ASSERT_TRUE(revoke.has_value()) << toString(revoke.error()); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + PartialProcessingExpirySettlesPendingCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + + auto start = + service.PutStart(client_id, "key", tenant_id, 100, MemoryConfig()); + ASSERT_TRUE(start.has_value()) << toString(start.error()); + AddCompletedDiskReplica(service, client_id, "key", tenant_id, 100); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 100); + + DiscardExpiredProcessingForTest(service, tenant_id, "key"); + + ExpectDiskOnlyObjectAndChargedBytes(service, tenant_id, "key", 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + DurablePartialProcessingExpirySettlesPendingCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + + auto start = + service.PutStart(client_id, "key", tenant_id, 100, MemoryConfig()); + ASSERT_TRUE(start.has_value()) << toString(start.error()); + AddCompletedDiskReplica(service, client_id, "key", tenant_id, 100); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 100); + + FinalizeExpiredProcessingForTest(service, tenant_id, "key"); + + ExpectDiskOnlyObjectAndChargedBytes(service, tenant_id, "key", 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + PartialSizeChangingUpsertRevokeReleasesReplacementCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + PutComplete(service, client_id, "key", tenant_id, 100); + + auto upsert = + service.UpsertStart(client_id, "key", tenant_id, 200, MemoryConfig()); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + AddCompletedDiskReplica(service, client_id, "key", tenant_id, 200); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 300); + + auto revoke = + service.UpsertRevoke(client_id, "key", tenant_id, ReplicaType::MEMORY); + + ASSERT_TRUE(revoke.has_value()) << toString(revoke.error()); + ExpectDiskOnlyObjectAndChargedBytes(service, tenant_id, "key", 0); +} + +TEST_F(MasterServiceTenantQuotaTest, + DurablePartialUpsertRevokeReleasesReplacementCharge) { + const TenantId tenant_id("tenant-a"); + MasterService service(MakeConfig({{tenant_id, 1000}})); + UUID client_id = MountSegment(service); + PutComplete(service, client_id, "key", tenant_id, 100); + + auto upsert = + service.UpsertStart(client_id, "key", tenant_id, 200, MemoryConfig()); + ASSERT_TRUE(upsert.has_value()) << toString(upsert.error()); + AddCompletedDiskReplica(service, client_id, "key", tenant_id, 200); + EXPECT_EQ(Snapshot(service, tenant_id).charged_bytes, 300); + + FinalizeRemovedMemoryReplicasForTest(service, tenant_id, "key"); + + ExpectDiskOnlyObjectAndChargedBytes(service, tenant_id, "key", 0); +} + TEST_F(MasterServiceTenantQuotaTest, CopyStartRequiresQuotaForNewReplica) { MasterService service(MakeConfig({{TenantId("tenant-a"), 150}})); UUID client_id = MountSegment(service, /*size=*/1024, "segment-a"); @@ -578,13 +857,10 @@ TEST_F(MasterServiceTenantQuotaTest, CopyStartRequiresQuotaForNewReplica) { ASSERT_FALSE(copy.has_value()); EXPECT_EQ(copy.error(), ErrorCode::TENANT_QUOTA_EXCEEDED); auto snapshot = Snapshot(service, TenantId("tenant-a")); - EXPECT_EQ(snapshot.used_bytes, 100); - EXPECT_EQ(snapshot.reserved_bytes, 0); - EXPECT_EQ(snapshot.committed_count, 1); + EXPECT_EQ(snapshot.charged_bytes, 100); } -TEST_F(MasterServiceTenantQuotaTest, - CopyEndCommitsAdditionalReplicaWithoutExtraObjectCount) { +TEST_F(MasterServiceTenantQuotaTest, CopyEndRetainsAdditionalReplicaCharge) { MasterService service(MakeConfig({{TenantId("tenant-a"), 300}})); UUID client_id = MountSegment(service, /*size=*/1024, "segment-a"); MountSegment(service, /*size=*/1024, "segment-b"); @@ -603,16 +879,38 @@ TEST_F(MasterServiceTenantQuotaTest, "segment-a", {"segment-b"}); ASSERT_TRUE(copy.has_value()) << toString(copy.error()); auto in_flight = Snapshot(service, TenantId("tenant-a")); - EXPECT_EQ(in_flight.used_bytes, 100); - EXPECT_EQ(in_flight.reserved_bytes, 100); + EXPECT_EQ(in_flight.charged_bytes, 200); ASSERT_TRUE( service.CopyEnd(client_id, "key", TenantId("tenant-a")).has_value()); auto completed = Snapshot(service, TenantId("tenant-a")); - EXPECT_EQ(completed.used_bytes, 200); - EXPECT_EQ(completed.reserved_bytes, 0); - EXPECT_EQ(completed.committed_count, 1); - EXPECT_EQ(completed.metadata_object_count, 1); + EXPECT_EQ(completed.charged_bytes, 200); +} + +TEST_F(MasterServiceTenantQuotaTest, CopyRevokeRefundsStartCharge) { + MasterService service(MakeConfig({{TenantId("tenant-a"), 300}})); + UUID client_id = MountSegment(service, /*size=*/1024, "segment-a"); + MountSegment(service, /*size=*/1024, "segment-b"); + + ReplicateConfig config = MemoryConfig(); + config.preferred_segment = "segment-a"; + auto put_start = + service.PutStart(client_id, "key", TenantId("tenant-a"), 100, config); + ASSERT_TRUE(put_start.has_value()) << toString(put_start.error()); + ASSERT_TRUE( + service + .PutEnd(client_id, "key", TenantId("tenant-a"), ReplicaType::MEMORY) + .has_value()); + + ASSERT_TRUE(service + .CopyStart(client_id, "key", TenantId("tenant-a"), + "segment-a", {"segment-b"}) + .has_value()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 200); + + ASSERT_TRUE( + service.CopyRevoke(client_id, "key", TenantId("tenant-a")).has_value()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 100); } TEST_F(MasterServiceTenantQuotaTest, @@ -637,8 +935,33 @@ TEST_F(MasterServiceTenantQuotaTest, ASSERT_FALSE(move.has_value()); EXPECT_EQ(move.error(), ErrorCode::TENANT_QUOTA_EXCEEDED); auto snapshot = Snapshot(service, TenantId("tenant-a")); - EXPECT_EQ(snapshot.used_bytes, 100); - EXPECT_EQ(snapshot.reserved_bytes, 0); + EXPECT_EQ(snapshot.charged_bytes, 100); +} + +TEST_F(MasterServiceTenantQuotaTest, MoveEndSettlesToFinalReplicaCharge) { + MasterService service(MakeConfig({{TenantId("tenant-a"), 300}})); + UUID client_id = MountSegment(service, /*size=*/1024, "segment-a"); + MountSegment(service, /*size=*/1024, "segment-b"); + + ReplicateConfig config = MemoryConfig(); + config.preferred_segment = "segment-a"; + ASSERT_TRUE( + service.PutStart(client_id, "key", TenantId("tenant-a"), 100, config) + .has_value()); + ASSERT_TRUE( + service + .PutEnd(client_id, "key", TenantId("tenant-a"), ReplicaType::MEMORY) + .has_value()); + + ASSERT_TRUE(service + .MoveStart(client_id, "key", TenantId("tenant-a"), + "segment-a", "segment-b") + .has_value()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 200); + + ASSERT_TRUE( + service.MoveEnd(client_id, "key", TenantId("tenant-a")).has_value()); + EXPECT_EQ(Snapshot(service, TenantId("tenant-a")).charged_bytes, 100); } TEST_F(MasterServiceTenantQuotaTest, DeletePolicyRequiresTenantWithoutObjects) { @@ -659,7 +982,7 @@ TEST_F(MasterServiceTenantQuotaTest, DeletePolicyRequiresTenantWithoutObjects) { } TEST_F(MasterServiceTenantQuotaTest, - DeletePolicyBlocksValidatedReservationsBeforeConnectorSave) { + DeletePolicyBlocksValidatedChargesBeforeConnectorSave) { MasterService service(MakeConfig({{TenantId("tenant-a"), 1000}})); MountSegment(service); @@ -686,14 +1009,14 @@ TEST_F(MasterServiceTenantQuotaTest, FAIL() << "timed out waiting for connector save"; } - auto reserve = ReserveTenantQuotaForTest(service, TenantId("tenant-a"), 1); - EXPECT_FALSE(reserve.has_value()); - EXPECT_EQ(reserve.error(), ErrorCode::TENANT_NOT_REGISTERED); + auto charge = ChargeTenantQuotaForTest(service, TenantId("tenant-a"), 1); + EXPECT_FALSE(charge.has_value()); + EXPECT_EQ(charge.error(), ErrorCode::TENANT_NOT_REGISTERED); - auto zero_byte_reserve = - ReserveTenantQuotaForTest(service, TenantId("tenant-a"), 0); - EXPECT_FALSE(zero_byte_reserve.has_value()); - EXPECT_EQ(zero_byte_reserve.error(), ErrorCode::TENANT_NOT_REGISTERED); + auto zero_byte_charge = + ChargeTenantQuotaForTest(service, TenantId("tenant-a"), 0); + EXPECT_FALSE(zero_byte_charge.has_value()); + EXPECT_EQ(zero_byte_charge.error(), ErrorCode::TENANT_NOT_REGISTERED); blocking_store_ptr->AllowSave(); delete_thread.join(); @@ -766,7 +1089,7 @@ TEST_F(MasterServiceTenantQuotaTest, #ifdef USE_NOF TEST_F(MasterServiceTenantQuotaTest, - DeletePolicyWaitsForZeroChargePutStartMetadataCreate) { + DeletePolicySeesZeroChargePutStartMetadataCreateWithoutPolicyLock) { MasterService service(MakeConfig({{TenantId("tenant-a"), 1000}})); UUID client_id = MountNoFSegment(service); @@ -781,6 +1104,7 @@ TEST_F(MasterServiceTenantQuotaTest, std::optional, ErrorCode>> put_result; + auto policy_lock = LockTenantQuotaPolicyForTest(service); std::thread put_thread([&] { put_result.emplace(service.PutStart(client_id, "nof-key", TenantId("tenant-a"), 128, config)); @@ -788,21 +1112,26 @@ TEST_F(MasterServiceTenantQuotaTest, if (allocation_started.wait_for(std::chrono::seconds(5)) != std::future_status::ready) { + policy_lock.unlock(); blocking_strategy_ptr->AllowAllocation(); put_thread.join(); - FAIL() << "timed out waiting for PutStart allocation"; + FAIL() << "zero-charge PutStart waited for tenant quota policy mutex"; } + policy_lock.unlock(); using DeleteResult = tl::expected, ErrorCode>; - std::optional delete_result; + std::promise delete_promise; + auto delete_future = delete_promise.get_future(); std::thread delete_thread([&] { - delete_result.emplace( + delete_promise.set_value( service.DeleteTenantQuotaPolicy(TenantId("tenant-a"))); }); - ASSERT_TRUE(WaitForTenantQuotaPolicyMutexContention(service)) - << "DeleteTenantQuotaPolicy did not wait for zero-charge PutStart"; + EXPECT_EQ(delete_future.wait_for(std::chrono::milliseconds(200)), + std::future_status::timeout) + << "tenant deletion passed the metadata scan while zero-charge " + "PutStart still held the target metadata shard"; blocking_strategy_ptr->AllowAllocation(); put_thread.join(); @@ -810,14 +1139,12 @@ TEST_F(MasterServiceTenantQuotaTest, ASSERT_TRUE(put_result.has_value()); ASSERT_TRUE(put_result->has_value()) << toString(put_result->error()); - ASSERT_TRUE(delete_result.has_value()); - ASSERT_FALSE(delete_result->has_value()); - EXPECT_EQ(delete_result->error(), ErrorCode::TENANT_NOT_EMPTY); + auto delete_result = delete_future.get(); + ASSERT_FALSE(delete_result.has_value()); + EXPECT_EQ(delete_result.error(), ErrorCode::TENANT_NOT_EMPTY); auto snapshot = Snapshot(service, TenantId("tenant-a")); - EXPECT_EQ(snapshot.used_bytes, 0); - EXPECT_EQ(snapshot.reserved_bytes, 0); - EXPECT_EQ(snapshot.metadata_object_count, 1); + EXPECT_EQ(snapshot.charged_bytes, 0); } #endif diff --git a/mooncake-store/tests/promotion_on_hit_test.cpp b/mooncake-store/tests/promotion_on_hit_test.cpp index dc14273626..2e86cd8da5 100644 --- a/mooncake-store/tests/promotion_on_hit_test.cpp +++ b/mooncake-store/tests/promotion_on_hit_test.cpp @@ -1075,6 +1075,10 @@ TEST_F(PromotionOnHitTest, HeartbeatBoundedBatchPreservesLeftovers) { TEST_F(PromotionOnHitTest, ReaperPopsStagedMemoryReplicaOnExpiry) { MasterServiceConfig config; config.enable_offload = true; + config.enable_multi_tenants = true; + config.tenant_quota_connector_type = "file"; + config.tenant_quota_connector_uri = + WriteTenantQuotaPolicyFile({{TenantId::Default().value(), 4096}}); config.promotion_on_hit = true; config.promotion_admission_threshold = 1; config.default_kv_lease_ttl = 2000; @@ -1103,6 +1107,9 @@ TEST_F(PromotionOnHitTest, ReaperPopsStagedMemoryReplicaOnExpiry) { auto alloc = service->PromotionAllocStart(ctx.client_id, "k_cold", TenantId::Default(), 1024, {}); ASSERT_TRUE(alloc.has_value()); + ASSERT_EQ( + service->GetTenantQuotaSnapshot(TenantId::Default())->charged_bytes, + 1024); // After AllocStart, the DRAM allocator must have committed bytes for // the staged PROCESSING MEMORY replica. @@ -1134,6 +1141,8 @@ TEST_F(PromotionOnHitTest, ReaperPopsStagedMemoryReplicaOnExpiry) { << "be freed back to the DRAM allocator. If this fires, the " << "reaper is not popping the staged replica and the buffer " << "leaks until the object itself is removed or evicted."; + EXPECT_EQ( + service->GetTenantQuotaSnapshot(TenantId::Default())->charged_bytes, 0); // NotifyPromotionSuccess for a reaped task must not commit anything // and must return REPLICA_IS_NOT_READY (the task entry is gone, so @@ -1386,6 +1395,89 @@ TEST_F(PromotionOnHitTest, NotifySuccessDecrementsCounter) { service->RemoveAll(); } +TEST_F(PromotionOnHitTest, TenantQuotaChargesAtAllocStartAndSettlesLifecycle) { + const TenantId tenant_id("tenant-a"); + MasterServiceConfig config; + config.enable_offload = true; + config.enable_multi_tenants = true; + config.tenant_quota_connector_type = "file"; + config.tenant_quota_connector_uri = + WriteTenantQuotaPolicyFile({{tenant_id.value(), 4096}}); + config.promotion_on_hit = true; + config.promotion_admission_threshold = 1; + config.default_kv_lease_ttl = 2000; + auto service = std::make_unique(config); + + constexpr size_t seg_size = 1024 * 1024 * 16; + auto seg = + PrepareSegment(*service, "seg_quota", kDefaultSegmentBase, seg_size); + ASSERT_TRUE(InjectLocalDiskReplica(*service, seg.client_id, "success", 1024, + seg.segment_name, tenant_id.value())); + ASSERT_TRUE(InjectLocalDiskReplica(*service, seg.client_id, "failure", 1024, + seg.segment_name, tenant_id.value())); + + ASSERT_TRUE(service->GetReplicaList("success", tenant_id)); + ASSERT_TRUE(service->PromotionAllocStart(seg.client_id, "success", + tenant_id, 1024, {})); + ASSERT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 1024); + ASSERT_TRUE( + service->NotifyPromotionSuccess(seg.client_id, "success", tenant_id)); + EXPECT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 1024); + + ASSERT_TRUE(service->GetReplicaList("failure", tenant_id)); + ASSERT_TRUE(service->PromotionAllocStart(seg.client_id, "failure", + tenant_id, 1024, {})); + ASSERT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 2048); + ASSERT_TRUE( + service->NotifyPromotionFailure(seg.client_id, "failure", tenant_id)); + EXPECT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 1024); + + ASSERT_TRUE(service->Remove("success", tenant_id, /*force=*/true)); + EXPECT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, 0); + service->RemoveAll(); +} + +TEST_F(PromotionOnHitTest, UpsertStartRejectsActivePromotionTask) { + const TenantId tenant_id("tenant-a"); + MasterServiceConfig config; + config.enable_offload = true; + config.enable_multi_tenants = true; + config.tenant_quota_connector_type = "file"; + config.tenant_quota_connector_uri = + WriteTenantQuotaPolicyFile({{tenant_id.value(), 4096}}); + config.promotion_on_hit = true; + config.promotion_admission_threshold = 1; + config.default_kv_lease_ttl = 2000; + auto service = std::make_unique(config); + + constexpr size_t seg_size = 1024 * 1024 * 16; + constexpr uint64_t object_size = 1024; + auto seg = + PrepareSegment(*service, "seg_quota", kDefaultSegmentBase, seg_size); + ASSERT_TRUE(InjectLocalDiskReplica(*service, seg.client_id, "key", + object_size, seg.segment_name, + tenant_id.value())); + + ASSERT_TRUE(service->GetReplicaList("key", tenant_id)); + ASSERT_TRUE(service->PromotionAllocStart(seg.client_id, "key", tenant_id, + object_size, {})); + ASSERT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, + object_size); + + ReplicateConfig replicate_config; + replicate_config.replica_num = 1; + auto upsert = service->UpsertStart(seg.client_id, "key", tenant_id, + object_size, replicate_config); + ASSERT_FALSE(upsert.has_value()); + EXPECT_EQ(upsert.error(), ErrorCode::OBJECT_HAS_REPLICATION_TASK); + + ASSERT_TRUE( + service->NotifyPromotionSuccess(seg.client_id, "key", tenant_id)); + EXPECT_EQ(service->GetTenantQuotaSnapshot(tenant_id)->charged_bytes, + object_size); + service->RemoveAll(); +} + // PromotionAllocStart must reject when the in-flight task has been // reaped between the holder's heartbeat and the AllocStart RPC arriving // (e.g. client stall past put_start_release_timeout_sec_). Without the diff --git a/mooncake-store/tests/tenant_quota_ledger_test.cpp b/mooncake-store/tests/tenant_quota_ledger_test.cpp new file mode 100644 index 0000000000..9ebdc1af02 --- /dev/null +++ b/mooncake-store/tests/tenant_quota_ledger_test.cpp @@ -0,0 +1,219 @@ +#include "tenant_quota_ledger.h" + +#include + +namespace mooncake { +namespace { + +class TenantQuotaLedgerTest : public ::testing::Test { + protected: + void SetUp() override { + ASSERT_TRUE(table_.UpsertTenantPolicy( + tenant_id_, TenantQuotaAccount::kMaxChargedBytes)); + table_.RecomputeEffectiveQuotas(TenantQuotaAccount::kMaxChargedBytes); + account_ = table_.GetOrCreateTenantHandle(tenant_id_); + } + + void Charge(uint64_t bytes) { ASSERT_TRUE(account_->TryCharge(bytes)); } + + const TenantId tenant_id_{"ledger-test"}; + TenantQuotaTable table_; + TenantQuotaHandle account_{nullptr}; +}; + +TEST_F(TenantQuotaLedgerTest, PendingFullyCommits) { + TenantQuotaLedger ledger; + Charge(100); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 100)); + + EXPECT_TRUE(ledger.SettlePrimaryWrite(account_, 100)); + EXPECT_EQ(ledger.PendingBytes(), 0); + EXPECT_EQ(ledger.CommittedBytes(), 100); + EXPECT_EQ(ledger.TotalChargedBytes(), 100); + EXPECT_EQ(account_->ChargedBytes(), 100); +} + +TEST_F(TenantQuotaLedgerTest, PendingPartiallyCommits) { + TenantQuotaLedger ledger; + Charge(100); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 100)); + + EXPECT_TRUE(ledger.SettlePrimaryWrite(account_, 60)); + EXPECT_EQ(ledger.PendingBytes(), 0); + EXPECT_EQ(ledger.CommittedBytes(), 60); + EXPECT_EQ(account_->ChargedBytes(), 60); +} + +TEST_F(TenantQuotaLedgerTest, PendingFullyRefunds) { + TenantQuotaLedger ledger; + Charge(100); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 100)); + + EXPECT_TRUE(ledger.RefundPending(account_)); + EXPECT_EQ(ledger.TotalChargedBytes(), 0); + EXPECT_EQ(account_->ChargedBytes(), 0); +} + +TEST_F(TenantQuotaLedgerTest, SettlesAdditionalCharge) { + TenantQuotaLedger ledger; + Charge(40); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 40)); + ASSERT_TRUE(ledger.SettlePrimaryWrite(account_, 40)); + Charge(60); + uint64_t task_pending_bytes = 60; + + EXPECT_TRUE(ledger.SettleAdditional(account_, task_pending_bytes, 35)); + EXPECT_EQ(task_pending_bytes, 60); + EXPECT_EQ(ledger.CommittedBytes(), 75); + EXPECT_EQ(account_->ChargedBytes(), 75); +} + +TEST_F(TenantQuotaLedgerTest, ReleasesCommittedBytesPartially) { + TenantQuotaLedger ledger; + Charge(100); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 100)); + ASSERT_TRUE(ledger.SettlePrimaryWrite(account_, 100)); + + EXPECT_TRUE(ledger.ReleaseCommitted(account_, 40)); + EXPECT_EQ(ledger.CommittedBytes(), 60); + EXPECT_EQ(account_->ChargedBytes(), 60); +} + +TEST_F(TenantQuotaLedgerTest, PrimaryWriteSettlementReleasesReplacementCharge) { + TenantQuotaLedger old_ledger; + TenantQuotaLedger replacement_owner; + TenantQuotaLedger new_ledger; + Charge(80); + ASSERT_TRUE(old_ledger.AdoptPendingCharge(account_, 80)); + ASSERT_TRUE(old_ledger.SettlePrimaryWrite(account_, 80)); + Charge(120); + ASSERT_TRUE(new_ledger.AdoptPendingCharge(account_, 120)); + + ASSERT_TRUE( + old_ledger.TransferReplacementCharge(account_, replacement_owner)); + ASSERT_TRUE( + replacement_owner.TransferReplacementCharge(account_, new_ledger)); + ASSERT_TRUE(new_ledger.SettlePrimaryWrite(account_, 100)); + EXPECT_EQ(new_ledger.PendingBytes(), 0); + EXPECT_EQ(new_ledger.CommittedBytes(), 100); + EXPECT_EQ(new_ledger.ReplacedBytes(), 0); + EXPECT_EQ(new_ledger.TotalChargedBytes(), 100); + EXPECT_EQ(account_->ChargedBytes(), 100); +} + +TEST_F(TenantQuotaLedgerTest, + PrimaryWriteMismatchPreservesPendingAndReplacementCharge) { + TenantQuotaLedger old_ledger; + TenantQuotaLedger new_ledger; + Charge(30); + ASSERT_TRUE(old_ledger.AdoptPendingCharge(account_, 30)); + ASSERT_TRUE(old_ledger.SettlePrimaryWrite(account_, 30)); + ASSERT_TRUE(old_ledger.TransferReplacementCharge(account_, new_ledger)); + Charge(70); + ASSERT_TRUE(new_ledger.AdoptPendingCharge(account_, 70)); + + auto settle_result = new_ledger.SettlePrimaryWrite(account_, 80); + + ASSERT_FALSE(settle_result); + EXPECT_EQ(settle_result.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(new_ledger.PendingBytes(), 70); + EXPECT_EQ(new_ledger.CommittedBytes(), 0); + EXPECT_EQ(new_ledger.ReplacedBytes(), 30); + EXPECT_EQ(new_ledger.TotalChargedBytes(), 100); + EXPECT_EQ(account_->ChargedBytes(), 100); + + ASSERT_TRUE(account_->Release(80)); + auto account_mismatch = new_ledger.SettlePrimaryWrite(account_, 50); + ASSERT_FALSE(account_mismatch); + EXPECT_EQ(account_mismatch.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(new_ledger.PendingBytes(), 70); + EXPECT_EQ(new_ledger.CommittedBytes(), 0); + EXPECT_EQ(new_ledger.ReplacedBytes(), 30); + EXPECT_EQ(new_ledger.TotalChargedBytes(), 100); + EXPECT_EQ(account_->ChargedBytes(), 20); +} + +TEST_F(TenantQuotaLedgerTest, RollsBackReplacementChargeOnFailure) { + TenantQuotaLedger old_ledger; + TenantQuotaLedger replacement_owner; + Charge(80); + ASSERT_TRUE(old_ledger.AdoptPendingCharge(account_, 80)); + ASSERT_TRUE(old_ledger.SettlePrimaryWrite(account_, 80)); + ASSERT_TRUE( + old_ledger.TransferReplacementCharge(account_, replacement_owner)); + + EXPECT_TRUE(replacement_owner.ReleaseReplacement(account_)); + EXPECT_EQ(replacement_owner.TotalChargedBytes(), 0); + EXPECT_EQ(account_->ChargedBytes(), 0); +} + +TEST_F(TenantQuotaLedgerTest, AccountingMismatchDoesNotMutateState) { + TenantQuotaLedger ledger; + Charge(50); + ASSERT_TRUE(ledger.AdoptPendingCharge(account_, 50)); + + auto settle_result = ledger.SettlePrimaryWrite(account_, 51); + ASSERT_FALSE(settle_result); + EXPECT_EQ(settle_result.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(ledger.PendingBytes(), 50); + EXPECT_EQ(account_->ChargedBytes(), 50); + + ASSERT_TRUE(ledger.SettlePrimaryWrite(account_, 50)); + auto release_result = ledger.ReleaseCommitted(account_, 51); + ASSERT_FALSE(release_result); + EXPECT_EQ(release_result.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(ledger.CommittedBytes(), 50); + EXPECT_EQ(account_->ChargedBytes(), 50); + + ASSERT_TRUE(ledger.ReleaseAll(account_)); + auto duplicate_release_all = ledger.ReleaseAll(account_); + ASSERT_FALSE(duplicate_release_all); + EXPECT_EQ(duplicate_release_all.error(), + TenantQuotaError::kAccountingMismatch); + + TenantQuotaLedger inconsistent_ledger; + ASSERT_TRUE(inconsistent_ledger.Rebuild(account_, 10)); + auto account_mismatch = inconsistent_ledger.ReleaseCommitted(account_, 10); + ASSERT_FALSE(account_mismatch); + EXPECT_EQ(account_mismatch.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(inconsistent_ledger.CommittedBytes(), 10); + EXPECT_EQ(account_->ChargedBytes(), 0); +} + +TEST_F(TenantQuotaLedgerTest, DuplicateReplacementOperationsAreRejected) { + TenantQuotaLedger source; + TenantQuotaLedger destination; + Charge(30); + ASSERT_TRUE(source.AdoptPendingCharge(account_, 30)); + ASSERT_TRUE(source.SettlePrimaryWrite(account_, 30)); + ASSERT_TRUE(source.TransferReplacementCharge(account_, destination)); + + auto duplicate_transfer = + source.TransferReplacementCharge(account_, destination); + ASSERT_FALSE(duplicate_transfer); + EXPECT_EQ(duplicate_transfer.error(), + TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(destination.ReplacedBytes(), 30); + EXPECT_EQ(account_->ChargedBytes(), 30); + + ASSERT_TRUE(destination.ReleaseReplacement(account_)); + auto duplicate_release = destination.ReleaseReplacement(account_); + ASSERT_FALSE(duplicate_release); + EXPECT_EQ(duplicate_release.error(), TenantQuotaError::kAccountingMismatch); + EXPECT_EQ(destination.TotalChargedBytes(), 0); + EXPECT_EQ(account_->ChargedBytes(), 0); +} + +TEST_F(TenantQuotaLedgerTest, RebuildMatchesGlobalChargedBytes) { + TenantQuotaLedger ledger; + ASSERT_TRUE(ledger.Rebuild(account_, 90)); + ASSERT_TRUE(table_.RebuildUsage({{tenant_id_, 90}})); + + EXPECT_EQ(ledger.PendingBytes(), 0); + EXPECT_EQ(ledger.CommittedBytes(), 90); + EXPECT_EQ(ledger.ReplacedBytes(), 0); + EXPECT_EQ(ledger.TotalChargedBytes(), account_->ChargedBytes()); +} + +} // namespace +} // namespace mooncake diff --git a/mooncake-store/tests/tenant_quota_test.cpp b/mooncake-store/tests/tenant_quota_test.cpp index e88b9d90e6..eca64238a0 100644 --- a/mooncake-store/tests/tenant_quota_test.cpp +++ b/mooncake-store/tests/tenant_quota_test.cpp @@ -105,9 +105,10 @@ void MakeOrphanTenant(TenantQuotaTable* table, const std::string& tenant_id, const TenantId canonical_tenant(tenant_id); ASSERT_TRUE(table->UpsertTenantPolicy(canonical_tenant, bytes).has_value()); table->RecomputeEffectiveQuotas(bytes); - ASSERT_TRUE(table->Reserve(canonical_tenant, bytes).has_value()); - ASSERT_TRUE(table->Commit(canonical_tenant, bytes).has_value()); - table->ApplyTenantPolicies({}); + ASSERT_TRUE(table->GetOrCreateTenantHandle(canonical_tenant) + ->TryCharge(bytes) + .has_value()); + ASSERT_TRUE(table->ApplyTenantPolicies({})); } TEST(TenantQuotaTableTest, NormalizesEmptyExplicitTenantIdToDefault) { @@ -137,6 +138,23 @@ TEST(TenantQuotaTableTest, RejectsZeroExplicitQuotaWithoutChangingState) { EXPECT_EQ(snapshot.requested_quota_bytes, 100); } +TEST(TenantQuotaTableTest, RejectsPolicyAboveAtomicAccountingRange) { + TenantQuotaTable table; + const TenantId tenant_id("tenant-a"); + ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 100)); + + auto upsert = table.UpsertTenantPolicy( + tenant_id, TenantQuotaAccount::kMaxChargedBytes + 1); + auto replace = table.ApplyTenantPolicies( + {{tenant_id, TenantQuotaAccount::kMaxChargedBytes + 1}}); + + ASSERT_FALSE(upsert); + EXPECT_EQ(upsert.error(), TenantQuotaError::kInvalidArgument); + ASSERT_FALSE(replace); + EXPECT_EQ(replace.error(), TenantQuotaError::kInvalidArgument); + EXPECT_EQ(Snapshot(table, "tenant-a").requested_quota_bytes, 100); +} + TEST(TenantQuotaTableTest, ApplyPoliciesCreatesOrphanState) { TenantQuotaTable table; MakeOrphanTenant(&table, "tenant-a", 40); @@ -146,18 +164,32 @@ TEST(TenantQuotaTableTest, ApplyPoliciesCreatesOrphanState) { EXPECT_FALSE(snapshot.has_explicit_policy); EXPECT_EQ(snapshot.requested_quota_bytes, 0); EXPECT_EQ(snapshot.effective_quota_bytes, 0); - EXPECT_EQ(snapshot.used_bytes, 40); - EXPECT_EQ(snapshot.committed_count, 1); + EXPECT_EQ(snapshot.charged_bytes, 40); + EXPECT_TRUE(snapshot.admission_closed); EXPECT_TRUE(snapshot.over_quota); } +TEST(TenantQuotaTableTest, ClosedOrphanAccountCanDrain) { + TenantQuotaTable table; + const TenantId tenant_id("tenant-a"); + MakeOrphanTenant(&table, tenant_id.value(), 40); + auto* handle = table.GetOrCreateTenantHandle(tenant_id); + + ASSERT_TRUE(handle->Release(40)); + + EXPECT_EQ(handle->ChargedBytes(), 0); + EXPECT_TRUE(handle->AdmissionClosed()); + EXPECT_FALSE(table.GetTenantSnapshot(tenant_id).has_value()); +} + TEST(TenantQuotaTableTest, ApplyPoliciesReplacesCanonicalPolicySet) { TenantQuotaTable table; const TenantId tenant_a("tenant-a"); const TenantId tenant_b("tenant-b"); const TenantId tenant_c("tenant-c"); table.ApplyTenantPolicies({{tenant_a, 100}, {tenant_b, 200}}); - table.IncrementMetadataObjectCount(tenant_a); + table.RecomputeEffectiveQuotas(300); + ASSERT_TRUE(table.GetOrCreateTenantHandle(tenant_a)->TryCharge(1)); table.ApplyTenantPolicies({{tenant_b, 300}, {tenant_c, 400}}); @@ -189,9 +221,11 @@ TEST(TenantQuotaTableTest, PolicyMutationDoesNotRecomputeEffectiveQuota) { ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 200).has_value()); EXPECT_EQ(Snapshot(table, "tenant-a").effective_quota_bytes, 100); + EXPECT_TRUE(Snapshot(table, "tenant-a").admission_closed); table.RecomputeEffectiveQuotas(1000); EXPECT_EQ(Snapshot(table, "tenant-a").effective_quota_bytes, 200); + EXPECT_FALSE(Snapshot(table, "tenant-a").admission_closed); } TEST(TenantQuotaTableTest, ListSnapshotsSortedAndCleansLazyEmptyTenants) { @@ -254,34 +288,32 @@ TEST(TenantQuotaTableTest, LazyEmptyOrphansDoNotAppearInList) { EXPECT_EQ(snapshots[0].tenant_id, TenantId("team-a")); } -TEST(TenantQuotaTableTest, ReserveRequiresRegisteredTenantIncludingZeroBytes) { +TEST(TenantQuotaTableTest, ChargeRequiresRegisteredTenantIncludingZeroBytes) { TenantQuotaTable table; + auto* account = table.GetOrCreateTenantHandle(TenantId("missing")); - auto regular = table.Reserve(TenantId("missing"), 1); - auto zero = table.Reserve(TenantId("missing"), 0); + auto regular = account->TryCharge(1); + auto zero = account->TryCharge(0); ASSERT_FALSE(regular.has_value()); ASSERT_FALSE(zero.has_value()); - EXPECT_EQ(regular.error(), TenantQuotaError::kTenantNotRegistered); - EXPECT_EQ(zero.error(), TenantQuotaError::kTenantNotRegistered); + EXPECT_EQ(regular.error().error, TenantQuotaError::kTenantNotRegistered); + EXPECT_EQ(zero.error().error, TenantQuotaError::kTenantNotRegistered); } -TEST(TenantQuotaTableTest, TracksAdditionalCommitAndMetadataCount) { +TEST(TenantQuotaTableTest, ChargeAndReleaseUpdateSingleCounter) { TenantQuotaTable table; const TenantId tenant_id("tenant-a"); ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 300).has_value()); table.RecomputeEffectiveQuotas(300); + auto* account = table.GetOrCreateTenantHandle(tenant_id); - ASSERT_TRUE(table.Reserve(tenant_id, 200).has_value()); - ASSERT_TRUE(table.Commit(tenant_id, 100).has_value()); - ASSERT_TRUE(table.CommitAdditional(tenant_id, 100).has_value()); - table.IncrementMetadataObjectCount(tenant_id); + ASSERT_TRUE(account->TryCharge(200).has_value()); + ASSERT_TRUE(account->Release(50).has_value()); auto snapshot = Snapshot(table, "tenant-a"); - EXPECT_EQ(snapshot.used_bytes, 200); - EXPECT_EQ(snapshot.reserved_bytes, 0); - EXPECT_EQ(snapshot.committed_count, 1); - EXPECT_EQ(snapshot.metadata_object_count, 1); + EXPECT_EQ(snapshot.charged_bytes, 150); + EXPECT_FALSE(snapshot.admission_closed); } TEST(TenantQuotaTableTest, AccountingMismatchDoesNotMutateState) { @@ -289,40 +321,13 @@ TEST(TenantQuotaTableTest, AccountingMismatchDoesNotMutateState) { const TenantId tenant_id("tenant-a"); ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 100).has_value()); table.RecomputeEffectiveQuotas(100); - ASSERT_TRUE(table.Reserve(tenant_id, 10).has_value()); - - auto commit = table.Commit(tenant_id, 11); - auto abort = table.Abort(tenant_id, 11); - auto before_commit = Snapshot(table, "tenant-a"); - ASSERT_FALSE(commit.has_value()); - ASSERT_FALSE(abort.has_value()); - EXPECT_EQ(commit.error(), TenantQuotaError::kAccountingMismatch); - EXPECT_EQ(abort.error(), TenantQuotaError::kAccountingMismatch); - EXPECT_EQ(before_commit.used_bytes, 0); - EXPECT_EQ(before_commit.reserved_bytes, 10); - - ASSERT_TRUE(table.Commit(tenant_id, 10).has_value()); - auto release = table.Release(tenant_id, 11); - auto partial = table.ReleasePartial(tenant_id, 11); + auto* account = table.GetOrCreateTenantHandle(tenant_id); + ASSERT_TRUE(account->TryCharge(10).has_value()); + + auto release = account->Release(11); ASSERT_FALSE(release.has_value()); - ASSERT_FALSE(partial.has_value()); EXPECT_EQ(release.error(), TenantQuotaError::kAccountingMismatch); - EXPECT_EQ(partial.error(), TenantQuotaError::kAccountingMismatch); - auto after_commit = Snapshot(table, "tenant-a"); - EXPECT_EQ(after_commit.used_bytes, 10); - EXPECT_EQ(after_commit.committed_count, 1); - - table.RebuildUsage({{tenant_id, - {.used_bytes = 10, - .committed_count = 0, - .metadata_object_count = 1}}}); - auto inconsistent_release = table.Release(tenant_id, 5); - ASSERT_FALSE(inconsistent_release.has_value()); - EXPECT_EQ(inconsistent_release.error(), - TenantQuotaError::kAccountingMismatch); - auto after_inconsistent_release = Snapshot(table, "tenant-a"); - EXPECT_EQ(after_inconsistent_release.used_bytes, 10); - EXPECT_EQ(after_inconsistent_release.committed_count, 0); + EXPECT_EQ(Snapshot(table, "tenant-a").charged_bytes, 10); } TEST(TenantQuotaTableTest, DisablePolicyRejectsNonEmptyTenant) { @@ -330,13 +335,36 @@ TEST(TenantQuotaTableTest, DisablePolicyRejectsNonEmptyTenant) { const TenantId tenant_id("tenant-a"); ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 100).has_value()); table.RecomputeEffectiveQuotas(100); - ASSERT_TRUE(table.Reserve(tenant_id, 1).has_value()); + ASSERT_TRUE( + table.GetOrCreateTenantHandle(tenant_id)->TryCharge(1).has_value()); auto result = table.DisableTenantPolicyIfEmpty(tenant_id); ASSERT_FALSE(result.has_value()); EXPECT_EQ(result.error(), TenantQuotaError::kTenantNotEmpty); EXPECT_TRUE(table.IsTenantRegistered(tenant_id)); + EXPECT_FALSE(Snapshot(table, "tenant-a").admission_closed); + EXPECT_TRUE(table.GetOrCreateTenantHandle(tenant_id)->TryCharge(1)); +} + +TEST(TenantQuotaTableTest, HandleRemainsStableAcrossPolicyLifecycle) { + TenantQuotaTable table; + const TenantId tenant_id("tenant-a"); + ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 100)); + table.RecomputeEffectiveQuotas(100); + auto* handle = table.GetOrCreateTenantHandle(tenant_id); + + ASSERT_TRUE(handle->TryCharge(10)); + ASSERT_TRUE(handle->Release(10)); + ASSERT_TRUE(table.DisableTenantPolicyIfEmpty(tenant_id)); + EXPECT_FALSE(handle->TryCharge(0)); + EXPECT_FALSE(table.GetTenantSnapshot(tenant_id).has_value()); + + ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, 200)); + table.RecomputeEffectiveQuotas(200); + EXPECT_EQ(table.GetOrCreateTenantHandle(tenant_id), handle); + EXPECT_TRUE(handle->TryCharge(200)); + EXPECT_EQ(Snapshot(table, "tenant-a").charged_bytes, 200); } TEST(TenantQuotaTableTest, RebuildUsageCreatesAndRemovesOrphans) { @@ -346,19 +374,17 @@ TEST(TenantQuotaTableTest, RebuildUsageCreatesAndRemovesOrphans) { ASSERT_TRUE(table.UpsertTenantPolicy(explicit_tenant, 100).has_value()); TenantQuotaUsageMap usage{ - {explicit_tenant, - {.used_bytes = 40, .committed_count = 1, .metadata_object_count = 1}}, - {orphan, - {.used_bytes = 20, .committed_count = 1, .metadata_object_count = 1}}, + {explicit_tenant, 40}, + {orphan, 20}, }; - table.RebuildUsage(usage); + ASSERT_TRUE(table.RebuildUsage(usage)); table.RecomputeEffectiveQuotas(100); EXPECT_TRUE(Snapshot(table, "tenant-a").has_explicit_policy); EXPECT_FALSE(Snapshot(table, "orphan").has_explicit_policy); EXPECT_TRUE(Snapshot(table, "orphan").over_quota); - table.RebuildUsage({}); + ASSERT_TRUE(table.RebuildUsage({})); EXPECT_FALSE(table.GetTenantSnapshot(orphan).has_value()); EXPECT_TRUE(table.GetTenantSnapshot(explicit_tenant).has_value()); } @@ -366,40 +392,34 @@ TEST(TenantQuotaTableTest, RebuildUsageCreatesAndRemovesOrphans) { TEST(TenantQuotaTableTest, OverflowChecksDoNotWrapAccounting) { TenantQuotaTable table; const TenantId tenant_id("tenant-a"); - const uint64_t max = std::numeric_limits::max(); + const uint64_t max = TenantQuotaAccount::kMaxChargedBytes; ASSERT_TRUE(table.UpsertTenantPolicy(tenant_id, max).has_value()); - table.RebuildUsage({{tenant_id, - {.used_bytes = max - 5, - .committed_count = max, - .metadata_object_count = max}}}); + ASSERT_TRUE(table.RebuildUsage({{tenant_id, max - 5}})); table.RecomputeEffectiveQuotas(max); - EXPECT_EQ(table.ComputeDeficit(tenant_id, 10), 5); - auto overflow_reserve = table.Reserve(tenant_id, 10); - ASSERT_FALSE(overflow_reserve.has_value()); - EXPECT_EQ(overflow_reserve.error(), TenantQuotaError::kQuotaExceeded); + auto* account = table.GetOrCreateTenantHandle(tenant_id); + auto overflow_charge = account->TryCharge(10); + ASSERT_FALSE(overflow_charge.has_value()); + EXPECT_EQ(overflow_charge.error().error, TenantQuotaError::kQuotaExceeded); + EXPECT_EQ(overflow_charge.error().deficit_bytes, 5); - ASSERT_TRUE(table.Reserve(tenant_id, 5).has_value()); - ASSERT_TRUE(table.Commit(tenant_id, 5).has_value()); - table.IncrementMetadataObjectCount(tenant_id); + ASSERT_TRUE(account->TryCharge(5).has_value()); auto snapshot = Snapshot(table, "tenant-a"); - EXPECT_EQ(snapshot.used_bytes, max); - EXPECT_EQ(snapshot.reserved_bytes, 0); - EXPECT_EQ(snapshot.committed_count, max); - EXPECT_EQ(snapshot.metadata_object_count, max); + EXPECT_EQ(snapshot.charged_bytes, max); } -TEST(ShardedTenantQuotaTableTest, ConcurrentReserveNeverExceedsQuota) { +TEST(ShardedTenantQuotaTableTest, ConcurrentChargeNeverExceedsQuota) { ShardedTenantQuotaTable<8> table; const TenantId tenant_id("tenant-a"); table.ApplyTenantPolicies({{tenant_id, 1000}}, 1000); + auto* account = table.GetOrCreateTenantHandle(tenant_id); std::atomic successes = 0; std::vector workers; for (int i = 0; i < 20; ++i) { workers.emplace_back([&] { - if (table.Reserve(tenant_id, 100).has_value()) { + if (account->TryCharge(100).has_value()) { ++successes; } }); @@ -409,7 +429,7 @@ TEST(ShardedTenantQuotaTableTest, ConcurrentReserveNeverExceedsQuota) { } EXPECT_EQ(successes.load(), 10); - EXPECT_EQ(Snapshot(table, "tenant-a").reserved_bytes, 1000); + EXPECT_EQ(Snapshot(table, "tenant-a").charged_bytes, 1000); } TEST(ShardedTenantQuotaTableTest, DifferentShardsUpdateIndependently) { @@ -424,50 +444,84 @@ TEST(ShardedTenantQuotaTableTest, DifferentShardsUpdateIndependently) { TestTable table; table.ApplyTenantPolicies({{tenant_a, 1000}, {tenant_b, 1000}}, 2000); + auto* account_a = table.GetOrCreateTenantHandle(tenant_a); + auto* account_b = table.GetOrCreateTenantHandle(tenant_b); std::atomic failures = 0; - auto update = [&](const TenantId& tenant_id) { + auto update = [&](TenantQuotaHandle account) { for (int i = 0; i < 1000; ++i) { - if (!table.Reserve(tenant_id, 1) || !table.Abort(tenant_id, 1)) { + if (!account->TryCharge(1) || !account->Release(1)) { ++failures; } } }; - std::thread first(update, std::cref(tenant_a)); - std::thread second(update, std::cref(tenant_b)); + std::thread first(update, account_a); + std::thread second(update, account_b); first.join(); second.join(); EXPECT_EQ(failures.load(), 0); - EXPECT_EQ(Snapshot(table, tenant_a.value()).reserved_bytes, 0); - EXPECT_EQ(Snapshot(table, tenant_b.value()).reserved_bytes, 0); + EXPECT_EQ(Snapshot(table, tenant_a.value()).charged_bytes, 0); + EXPECT_EQ(Snapshot(table, tenant_b.value()).charged_bytes, 0); +} + +TEST(ShardedTenantQuotaTableTest, + CrossShardMutationsValidateEverythingBeforeUpdating) { + using TestTable = ShardedTenantQuotaTable<2>; + const TenantId tenant_a("tenant-a"); + TenantId tenant_b("tenant-b"); + for (int suffix = 0; TenantIdHash{}(tenant_a) % TestTable::kNumShards == + TenantIdHash{}(tenant_b) % TestTable::kNumShards; + ++suffix) { + tenant_b = TenantId("tenant-b-" + std::to_string(suffix)); + } + + TestTable table; + ASSERT_TRUE( + table.ApplyTenantPolicies({{tenant_a, 100}, {tenant_b, 200}}, 300)); + auto invalid_policy = table.ApplyTenantPolicies( + {{tenant_a, 300}, {tenant_b, TenantQuotaAccount::kMaxChargedBytes + 1}}, + 500); + ASSERT_FALSE(invalid_policy); + EXPECT_EQ(Snapshot(table, tenant_a.value()).requested_quota_bytes, 100); + EXPECT_EQ(Snapshot(table, tenant_b.value()).requested_quota_bytes, 200); + + ASSERT_TRUE(table.RebuildUsage({{tenant_a, 40}, {tenant_b, 50}}, 300)); + auto invalid_usage = table.RebuildUsage( + {{tenant_a, 60}, {tenant_b, TenantQuotaAccount::kMaxChargedBytes + 1}}, + 300); + ASSERT_FALSE(invalid_usage); + EXPECT_EQ(Snapshot(table, tenant_a.value()).charged_bytes, 40); + EXPECT_EQ(Snapshot(table, tenant_b.value()).charged_bytes, 50); } TEST(ShardedTenantQuotaTableTest, - DisabledPolicyRejectsRegularAndZeroByteReservations) { + DisabledPolicyRejectsRegularAndZeroByteCharges) { ShardedTenantQuotaTable<8> table; const TenantId tenant_id("tenant-a"); table.ApplyTenantPolicies({{tenant_id, 100}}, 100); + auto* account = table.GetOrCreateTenantHandle(tenant_id); ASSERT_TRUE(table.DisableTenantPolicyIfEmpty(tenant_id).has_value()); - auto regular = table.Reserve(tenant_id, 1); - auto zero = table.Reserve(tenant_id, 0); + auto regular = account->TryCharge(1); + auto zero = account->TryCharge(0); ASSERT_FALSE(regular.has_value()); ASSERT_FALSE(zero.has_value()); - EXPECT_EQ(regular.error(), TenantQuotaError::kTenantNotRegistered); - EXPECT_EQ(zero.error(), TenantQuotaError::kTenantNotRegistered); + EXPECT_EQ(regular.error().error, TenantQuotaError::kTenantNotRegistered); + EXPECT_EQ(zero.error().error, TenantQuotaError::kTenantNotRegistered); } TEST(ShardedTenantQuotaTableTest, RecomputeCanRunWithAccounting) { ShardedTenantQuotaTable<8> table; const TenantId tenant_id("tenant-a"); table.ApplyTenantPolicies({{tenant_id, 1000}}, 1000); + auto* account = table.GetOrCreateTenantHandle(tenant_id); std::atomic failures = 0; std::thread accounting([&] { for (int i = 0; i < 1000; ++i) { - if (!table.Reserve(tenant_id, 1) || !table.Abort(tenant_id, 1)) { + if (!account->TryCharge(1) || !account->Release(1)) { ++failures; } } @@ -483,7 +537,7 @@ TEST(ShardedTenantQuotaTableTest, RecomputeCanRunWithAccounting) { EXPECT_EQ(failures.load(), 0); auto snapshot = Snapshot(table, "tenant-a"); - EXPECT_EQ(snapshot.reserved_bytes, 0); + EXPECT_EQ(snapshot.charged_bytes, 0); EXPECT_EQ(snapshot.effective_quota_bytes, 1000); } @@ -524,6 +578,8 @@ TEST(TenantQuotaPolicyStoreTest, RejectsInvalidYamlPolicies) { "version: 1\n\ntenants:\n - name: tenant-a\n quota: 1KB\n - name: " "tenant-a\n quota: 2KB\n", "version: 1\n\ntenants:\n - name: tenant-a\n quota: " + "9223372036854775808\n", + "version: 1\n\ntenants:\n - name: tenant-a\n quota: " "18446744073709551616\n", "version: 1\n\ntenants:\n - name: tenant-a\n quota: " "18446744073709551615TB\n", From 6cc13f0ed484d8bea87bda65be0c5fa996ec8dd3 Mon Sep 17 00:00:00 2001 From: Chizheng Fang <93508110+fcczzz@users.noreply.github.com> Date: Thu, 13 Aug 2026 10:31:19 +0800 Subject: [PATCH 045/483] [Store] Add DFS replica support with POSIX backend (#2683) Add DFS replica metadata, descriptor caching, shard-based global allocation, and POSIX-backed distributed storage IO. Wire DFS offload/read flows through master/client services and add POSIX DFS coverage. --- .../api-reference/cpp/mooncake-store.md | 34 +- .../api-reference/python/mooncake-store.md | 54 +- .../mooncake-store-deployment-guide.md | 212 ++- docs/source/design/mooncake-store.md | 88 +- .../plugin-usage/3FS-USRBIO-Plugin.md | 87 +- mooncake-common/include/environ.h | 3 +- mooncake-common/src/environ.cpp | 30 + mooncake-common/tests/environ_test.cpp | 16 + mooncake-integration/store/store_py.cpp | 1 + .../store/store_py_internal.h | 3 +- mooncake-store/include/allocator.h | 11 +- mooncake-store/include/client_service.h | 13 + mooncake-store/include/hf3fs/hf3fs.h | 6 +- mooncake-store/include/master_service.h | 13 +- mooncake-store/include/replica.h | 95 +- mooncake-store/include/replica_selection.h | 13 +- .../distributed/dfs_global_allocator.h | 166 +++ .../distributed/distributed_storage_backend.h | 51 +- .../include/storage/distributed/fs_adapter.h | 28 + .../storage/distributed/hf3fs_adapter.h | 14 + .../storage/distributed/posix_fs_adapter.h | 55 + mooncake-store/include/storage_backend.h | 3 + mooncake-store/include/types.h | 3 +- mooncake-store/include/utils.h | 1 + mooncake-store/src/CMakeLists.txt | 2 + mooncake-store/src/client_buffer.cpp | 5 + mooncake-store/src/client_service.cpp | 393 +++++- mooncake-store/src/file_storage.cpp | 46 +- mooncake-store/src/hf3fs/README.md | 79 +- mooncake-store/src/master_service.cpp | 356 ++++- mooncake-store/src/real_client.cpp | 192 ++- mooncake-store/src/serialize/serializer.cpp | 38 + .../distributed/dfs_global_allocator.cpp | 474 +++++++ .../distributed_storage_backend.cpp | 530 +++++--- .../src/storage/distributed/hf3fs_adapter.cpp | 194 ++- .../storage/distributed/posix_fs_adapter.cpp | 236 ++++ mooncake-store/src/storage_backend.cpp | 8 +- mooncake-store/src/types.cpp | 1 + mooncake-store/tests/CMakeLists.txt | 10 + mooncake-store/tests/dfs_hf3fs_test.cpp | 178 +++ mooncake-store/tests/dfs_posix_test.cpp | 1166 +++++++++++++++++ mooncake-store/tests/dfs_sync_client_test.cpp | 363 +++++ mooncake-store/tests/master_scenario.cpp | 22 +- mooncake-store/tests/master_scenario.h | 12 +- .../master_service_evict_scenario_test.cpp | 16 +- mooncake-store/tests/master_service_test.cpp | 196 +++ .../tests/replica_selection_test.cpp | 30 + 47 files changed, 5024 insertions(+), 523 deletions(-) create mode 100644 mooncake-store/include/storage/distributed/dfs_global_allocator.h create mode 100644 mooncake-store/include/storage/distributed/posix_fs_adapter.h create mode 100644 mooncake-store/src/storage/distributed/dfs_global_allocator.cpp create mode 100644 mooncake-store/src/storage/distributed/posix_fs_adapter.cpp create mode 100644 mooncake-store/tests/dfs_hf3fs_test.cpp create mode 100644 mooncake-store/tests/dfs_posix_test.cpp create mode 100644 mooncake-store/tests/dfs_sync_client_test.cpp diff --git a/docs/source/api-reference/cpp/mooncake-store.md b/docs/source/api-reference/cpp/mooncake-store.md index aca0e97953..49d9ca4ed8 100644 --- a/docs/source/api-reference/cpp/mooncake-store.md +++ b/docs/source/api-reference/cpp/mooncake-store.md @@ -26,7 +26,7 @@ tl::expected Get(const std::string& object_key, std::vector& slices); ``` -`Get` retrieves the value of `object_key` into the provided `slices`. The returned data is guaranteed to be complete and correct. Each slice must reference local DRAM/VRAM memory that has been pre-registered with `registerLocalMemory(addr, len)` (not the global segments that contribute to the distributed memory pool). When persistence is enabled and the requested data is not found in the distributed memory pool, `Get` will fall back to loading the data from SSD. +`Get` retrieves the value of `object_key` into the provided `slices`. The returned data is guaranteed to be complete and correct. Each slice must reference local DRAM/VRAM memory that has been pre-registered with `registerLocalMemory(addr, len)` (not the global segments that contribute to the distributed memory pool). The master returns the readable replica list and the client selects a complete replica. Depending on the selected replica, the data may be read from memory, NoF SSD, legacy shared-filesystem `DISK`, client-owned `LOCAL_DISK`, or the configured descriptor-based DFS backend. ### Put @@ -36,25 +36,36 @@ tl::expected Put(const ObjectKey& key, const ReplicateConfig& config); ``` -`Put` stores the value associated with `key` in the distributed memory pool. The `config` parameter allows specifying the required number of replicas as well as the preferred segment for storing the value. When persistence is enabled, `Put` also asynchronously triggers a persistence operation to SSD. +`Put` stores the value associated with `key` in the configured replica tiers. The `config` parameter controls the number of memory, NoF, and DFS replicas as well as placement preferences. Legacy `DISK` persistence and client-owned `LOCAL_DISK` SSD offload remain asynchronous. When `dfs_replica_num` is `1`, `Put` waits for the requested DFS `WriteAt` operation before returning success; this does not provide an additional `fsync` durability guarantee. -**Replication Guarantees and Best Effort Behavior:** +**Memory Replication Guarantees and Best Effort Behavior:** - Each slice of an object is guaranteed to be replicated to different segments, ensuring distribution across separate storage nodes - Different slices from different objects may be placed in the same segment - Replication operates on a best-effort basis: if insufficient space is available for all requested replicas, the object will still be written with as many replicas as possible -The data structure details of `ReplicateConfig` are as follows: +Requests with `dfs_replica_num == 1` use reliable multi-replica mode: allocation and every requested transfer must succeed, otherwise `Put` fails and allocated replicas are revoked. + +```{warning} +Descriptor-based DFS is a work-in-progress feature for development and evaluation. It is not covered by the Store's production fault-tolerance, HA, durability, or multi-tenant guarantees. +``` + +The DFS-related replica-count fields of `ReplicateConfig` are as follows: ```C++ struct ReplicateConfig { - size_t replica_num{1}; // Total number of replicas for the object + size_t replica_num{1}; // Memory replicas + size_t nof_replica_num{0}; // NoF SSD replicas + size_t dfs_replica_num{0}; // Shared DFS replicas (0 or 1) SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; - std::optional soft_pin_ttl_ms{}; // ENABLE override; omitted uses the Master default - bool with_hard_pin{false}; // Whether to enable hard pin (never evicted) - std::string preferred_segment{}; // Preferred segment for allocation + std::optional soft_pin_ttl_ms{}; // ENABLE override; omitted uses the Master default + bool with_hard_pin{false}; // Whether to enable hard pin (never evicted) + std::string preferred_segment{}; // Preferred segment for allocation + // Other placement, data-type, and grouping fields are omitted. }; ``` +`dfs_replica_num` may currently be `0` or `1`. When it is `1`, `replica_num >= 1` is required, so DFS-only placement is not supported. DFS replicas currently support only the `default` tenant and require the master and client DFS backends to be configured with the same shared root and shard layout. See the {ref}`DFS deployment documentation ` for setup and lifecycle limitations. Native C++ clients must initialize a `DistributedStorageBackend` and attach it with `SetDfsStorageBackend()` before issuing DFS reads or writes. DFS descriptors are carried by each `PutStart`, `UpsertStart`, or query response; there is no client-side descriptor cache. The Python/RealClient setup path attaches the backend through `FileStorage`. + Soft pinning starts when the first replica becomes readable and has a fixed lifetime: reads do not extend it. `PRESERVE` keeps the committed deadline on an Upsert, `ENABLE` starts a new lifetime, and `DISABLE` removes it when the write @@ -80,8 +91,11 @@ std::vector> BatchUpsert( `Upsert` inserts `key` if it does not exist and updates the existing object if it does. It uses the same replication configuration model as `Put`, while allowing the store to reuse existing placement for in-place updates when the -current layout permits it. `BatchUpsert` performs the same operation for -multiple keys using a shared replication configuration. +current layout permits it. If either the existing object or the new request has +a DFS replica, a same-size update requires the requested memory, NoF, and DFS +replica counts to match the existing topology. A different-size update releases +the old placement and allocates a new topology. `BatchUpsert` performs the same +operation for multiple keys using a shared replication configuration. ### Remove diff --git a/docs/source/api-reference/python/mooncake-store.md b/docs/source/api-reference/python/mooncake-store.md index bbce3845e7..48fdc075ab 100644 --- a/docs/source/api-reference/python/mooncake-store.md +++ b/docs/source/api-reference/python/mooncake-store.md @@ -602,13 +602,52 @@ config = ReplicateConfig() #### replica_num **Type:** `int` **Default:** `1` -**Description:** Specifies the total number of replicas to create for the stored object. +**Description:** Specifies the number of memory replicas to create for the +stored object. ```python config = ReplicateConfig() -config.replica_num = 3 # Store 3 copies of the data +config.replica_num = 3 # Store 3 memory replicas ``` +#### nof_replica_num +**Type:** `int` +**Default:** `0` +**Description:** Specifies the number of replicas to create in the configured +NVMe-oF SSD pool. + +```python +config = ReplicateConfig() +config.replica_num = 1 +config.nof_replica_num = 1 +``` + +#### dfs_replica_num +**Type:** `int` +**Default:** `0` +**Status:** **Work in progress; development and evaluation only.** +**Description:** Requests an additional replica in the configured shared +distributed filesystem. The supported values are currently `0` and `1`. When +set to `1`, `replica_num` must be at least `1`, so DFS-only placement is not +supported. DFS replicas currently support only the `default` tenant. + +```python +config = ReplicateConfig() +config.replica_num = 1 +config.dfs_replica_num = 1 +``` + +Writes that request a DFS replica return success after the DFS `WriteAt` +operation completes, but without an additional `fsync` durability guarantee. +The master and client DFS backends must be enabled and configured with the same +absolute shared-root path and shard layout. See the +{ref}`DFS deployment documentation ` for the required environment +variables and current limitations. + +For a same-size `upsert`, if either the existing object or the new request has +a DFS replica, the requested memory, NoF, and DFS replica counts must match the +existing topology. A different-size update allocates a new topology. + #### soft_pin_action **Type:** `SoftPinAction` **Default:** `SoftPinAction.PRESERVE` @@ -1101,8 +1140,15 @@ the positional overload. - `rdma_devices` (str): **Required by the positional overload**. RDMA/EFA device name(s), e.g. `"mlx5_0"` or `"mlx5_0,mlx5_1"`. Leave empty to auto-discover NICs unless `MC_MS_AUTO_DISC=0`; always empty for TCP. - `master_server_addr` (str): **Required by the positional overload**. Master server address (e.g., "localhost:50051") - `engine` (Optional[TransferEngine]): Existing Transfer Engine instance to reuse. Defaults to `None`. -- `enable_ssd_offload` (bool): Enable client-side SSD offload support. Defaults to `False`. -- `ssd_offload_path` (str): SSD offload directory. When provided, overrides the storage path environment configuration. +- `enable_ssd_offload` (bool): Initialize client-side `FileStorage`. With a + normal file backend this enables SSD offload; with + `MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR=distributed_storage_backend`, + it initializes the DFS backend and is required for DFS reads and writes. + Defaults to `False`. +- `ssd_offload_path` (str): FileStorage directory. When provided, it overrides + `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH`. With the distributed backend, DFS shard + data is stored under `MOONCAKE_DFS_ROOT_DIR`, but this separate directory is + still validated during FileStorage initialization. - `tenant_id` (str): Tenant namespace for object keys. Defaults to `"default"`. - `enable_client_http_server` (bool): Enable the client-local `/health`, `/metrics`, and `/metrics/summary` HTTP endpoints. Defaults to `False`. - `client_http_port` (int): Port for the client-local HTTP endpoints. Defaults to `9300`. diff --git a/docs/source/deployment/mooncake-store-deployment-guide.md b/docs/source/deployment/mooncake-store-deployment-guide.md index a269628a96..2c8f099e85 100644 --- a/docs/source/deployment/mooncake-store-deployment-guide.md +++ b/docs/source/deployment/mooncake-store-deployment-guide.md @@ -415,7 +415,7 @@ Multi-Tenant Deployment - Use `/metrics/summary` during bring-up; integrate `/metrics` with Prometheus/Grafana for production. - For detailed SSD offload configuration (storage backends, eviction policies, io_uring), see the [SSD Offload guide](ssd/ssd-offload). - For NVMe-oF SSD pool configuration see the [NVMe-oF SSD Pool Deployment Guide](ssd/nvmf-ssd-deployment-guide) -- For experimental 3FS (USRBIO) integration as a persistent storage backend, see the [3FS USRBIO Plugin guide](../getting_started/plugin-usage/3FS-USRBIO-Plugin). +- For the experimental HF3FS USRBIO adapter used by descriptor-based DFS replicas, see the [HF3FS USRBIO adapter guide](../getting_started/plugin-usage/3FS-USRBIO-Plugin). - For detailed monitoring and observation see [Observability](../getting_started/observability) :::{toctree} @@ -424,7 +424,7 @@ Multi-Tenant Deployment KV Cache Sharing and Isolation SSD Storage -HF3FS Plugin (Experimental)<../getting_started/plugin-usage/3FS-USRBIO-Plugin> +HF3FS USRBIO Adapter (Experimental)<../getting_started/plugin-usage/3FS-USRBIO-Plugin> ../getting_started/observability ::: @@ -676,14 +676,202 @@ When `--offload_on_evict=true` is active, each `BatchEvict` cycle can queue at m When `--allocation_strategy=cxl` is set alongside `--enable_cxl=true`, the master preferentially allocates new objects on CXL memory. -### DFS Storage +### Legacy Shared-filesystem `DISK` Persistence + +The older shared-filesystem persistence path remains available independently +of descriptor-based DFS: | Flag | Default | Description | |------|---------|-------------| -| `--root_fs_dir` | empty | Legacy DFS persistence directory; do not use with SSD offload | -| `--global_file_segment_size` | `INT64_MAX` (unlimited) | Max available space for DFS segments; default does not cap DFS usage | +| `--root_fs_dir` | empty | Enable legacy `DISK` replicas under `/`. The path must resolve to the same shared filesystem location on every participating client. | +| `--global_file_segment_size` | `INT64_MAX` (unlimited) | Declared legacy file capacity used by master usage metrics. It does not configure descriptor-based DFS shard files. | + +With `--root_fs_dir` set, the master adds a legacy `DISK` replica to each new +object and clients write it asynchronously. This path is distinct from both +client-owned `LOCAL_DISK` SSD offload and descriptor-based `DFS` replicas. Do +not combine `--root_fs_dir` with `--enable_offload=true`; configure real-client +SSD offload with `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` instead. + +(dfs-storage)= +### Descriptor-based DFS Storage + +```{warning} +**Work in progress.** Descriptor-based DFS is intended for development and +evaluation only. It is not production-ready and is not covered by Mooncake +Store's general fault-tolerance, HA continuity, durability, or multi-tenant +guarantees. +``` + +Mooncake Store can place an additional replica in a shared distributed +filesystem. The master allocates aligned ranges in pre-created shard files and +publishes a descriptor containing the shard, offset, and object size. Clients +use that descriptor to access the same files through either regular POSIX I/O +or the HF3FS USRBIO adapter. + +DFS replicas are separate from `LOCAL_DISK` SSD-offload replicas. They do not +use the legacy `--root_fs_dir` persistence path or the master's asynchronous +offload task queue. + +```{note} +DFS allocator state is not yet restored after a master restart or HA leader +failover. Do not enable descriptor-based DFS in a deployment that requires +master recovery, HA continuity, or multiple tenants. See the complete list of +limitations below. +``` + +#### Master configuration + +Enable the DFS allocator in the master process and select a shared root and +shard layout. For example, to use HF3FS: + +```bash +export MOONCAKE_ENABLE_DFS=1 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true + +mooncake_master [other master arguments] +``` + +At startup, the master creates `MOONCAKE_DFS_SHARD_COUNT` shard files and +preallocates each file to `MOONCAKE_DFS_SHARD_CAPACITY`. The example therefore +configures 256 GiB of total logical shard capacity (`64 * 4 GiB`). Ensure the +shared filesystem has sufficient capacity; whether all backing space is +reserved immediately depends on the selected filesystem adapter. + +The `hf3fs` adapter requires Mooncake to be built with `USE_3FS=ON`. Use +`MOONCAKE_DFS_FS_ADAPTER=posix` for development and integration testing on a +regular shared filesystem. + +#### Client configuration + +Every client that may read or write a DFS replica must initialize +`FileStorage` and select the distributed backend. Use an absolute DFS root path; +the root string, shard count, shard capacity, and alignment must match the +master configuration. Select an adapter that can access the same underlying +shared files; the examples use the same adapter in every process. + +```bash +export MOONCAKE_OFFLOAD_ENABLED=true +export MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR=distributed_storage_backend +export MOONCAKE_OFFLOAD_FILE_STORAGE_PATH=/data/file_storage +export MOONCAKE_MASTER=127.0.0.1:50051 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true + +python -m mooncake.mooncake_store_service +``` + +For a programmatic Python client, pass `enable_ssd_offload=True` to `setup()` +instead of `MOONCAKE_OFFLOAD_ENABLED`. Programmatic setup still reads the +backend-specific `MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR` and +`MOONCAKE_DFS_*` variables shown above; only the launcher-level setup fields are +supplied as Python arguments. The +`MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` directory must already exist and be an +absolute, writable, non-symlink directory. DFS shard data is stored under +`MOONCAKE_DFS_ROOT_DIR`; the FileStorage path is still required for client +initialization because the shared `FileStorageConfig` validates it even when +the selected backend stores data in the DFS root. + +Native C++ clients must initialize a `DistributedStorageBackend` with the same +DFS layout and attach it to the client with `SetDfsStorageBackend()` before +issuing DFS reads or writes. Reads and writes use the DFS descriptor carried by +the current query or start-operation response; no client-side descriptor cache +is required. + +#### DFS configuration reference + +| Variable | Scope | Default | Description | +|----------|-------|---------|-------------| +| `MOONCAKE_ENABLE_DFS` | Master | `false` | Enable master-side DFS allocation. `MOONCAKE_DFS_ENABLED` is accepted as a compatibility fallback. | +| `MOONCAKE_DFS_ROOT_DIR` | Master and clients | `/mnt/3fs/mooncake` | Absolute shared shard root; use the same path string in every process. Falls back to `MOONCAKE_DISTRIBUTED_ROOT_DIR`. | +| `MOONCAKE_DFS_FS_ADAPTER` | Master and clients | `hf3fs` | Filesystem adapter: `hf3fs` or `posix`. Falls back to `MOONCAKE_DISTRIBUTED_FS_TYPE`. | +| `MOONCAKE_DFS_SHARD_COUNT` | Master and clients | `64` | Number of DFS shard files. | +| `MOONCAKE_DFS_SHARD_CAPACITY` | Master and clients | `4294967296` (4 GiB) | Logical file capacity of each shard in bytes. Each object is allocated wholly within one shard. | +| `MOONCAKE_DFS_ALIGNMENT` | Master and clients | `4096` | Allocation alignment in bytes; must be a power of two and divide the shard capacity. | +| `MOONCAKE_DFS_SINGLE_TENANT` | Master and clients | `true` | Currently must remain `true`. | +| `MOONCAKE_DFS_EVICTION_ENABLED` | Master | `true` | Enable DFS allocator eviction. | +| `MOONCAKE_DFS_EVICTION_HIGH_WATERMARK` | Master | `0.9` | Usage ratio that triggers eviction. | +| `MOONCAKE_DFS_EVICTION_LOW_WATERMARK` | Master | `0.7` | Usage ratio targeted by an eviction cycle. | +| `MOONCAKE_DFS_DEFERRED_FREE_SECONDS` | Master | `30` | Delay before a freed shard range may be reused. | +| `MOONCAKE_DFS_EVICTION_CHECK_INTERVAL` | Master | `5` | Eviction check interval in seconds. | + +#### Requesting and accessing DFS replicas + +Callers request DFS placement through `ReplicateConfig`: + +```python +from mooncake.store import ReplicateConfig + +config = ReplicateConfig() +config.replica_num = 1 +config.dfs_replica_num = 1 +store.put("key", b"value", config) +``` -`--root_fs_dir` is a legacy persistence parameter and is expected to be replaced as the distributed filesystem path is refactored. For SSD offload, configure `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` on each real client instead. +`dfs_replica_num` may currently be `0` or `1`. A DFS replica must be requested +with at least one memory replica (`replica_num >= 1`), so DFS-only placement is +not supported. + +Each key hashes to exactly one DFS shard. Allocation does not fall back to a +different shard, so a request may return `NO_AVAILABLE_HANDLE` when its selected +shard is full even if other shards have free space. A DFS object is never +striped across shards. The selected shard must have room for the object rounded +up to `MOONCAKE_DFS_ALIGNMENT`, plus up to one alignment unit of allocator +padding (`MOONCAKE_DFS_ALIGNMENT - 1` bytes); usable object capacity is +therefore lower than the shard file's +logical size. + +For `Put`, `BatchPut`, `Upsert`, and `BatchUpsert`, the client writes requested +memory and NoF replicas, stages device buffers to host memory when necessary, +and then performs positional DFS writes. A successful request means the +requested DFS `WriteAt` operations completed. It does **not** imply that an +additional `fsync` completed. Batch operations isolate failures by key; a +failed key is revoked without downgrading successful keys. + +For a same-size `Upsert`, if either the existing object or the new request has +a DFS replica, the requested memory, NoF, and DFS replica counts must match the +existing topology. A different-size update releases the old placement and +allocates a new topology. + +On reads, the master returns the readable replica list through the normal query +path, and the client selects the first complete replica. If it selects DFS, any +client configured with the same DFS root and shard layout can issue positional +reads for that descriptor. + +#### Current limitations + +- Only the `default` tenant is supported. +- `dfs_replica_num` must be `0` or `1`, and `replica_num >= 1` is required when + it is enabled. +- C and Rust clients cannot currently request or access descriptor-based DFS: + their replication configuration does not expose `dfs_replica_num`, and their + setup API cannot initialize the distributed `FileStorage` backend. Use the + native C++ or Python/RealClient API. +- A DFS object must fit in its key-selected shard after alignment and allocator + padding; objects are not striped and allocation does not fall back to another + shard. +- DFS allocator state is currently in memory. A master restart or HA leader + failover does not reconstruct existing DFS allocations, so DFS cannot provide + continuity across those events. +- DFS cannot be enabled with snapshot generation, snapshot restore, oplog + recovery, or standby restore until DFS allocator state restoration is + implemented. +- There is currently no background DFS retry queue or configurable + asynchronous acknowledgement policy. +- DFS writes currently have no DFS-specific timeout, request cancellation, or + `fsync` durability guarantee. + +The older `--root_fs_dir` and `--global_file_segment_size` flags configure the +legacy `DISK` path described above and are not used by descriptor-based DFS +replicas. ### NoF (NVMe-oF SSD Pool) @@ -735,7 +923,7 @@ The client derives the host id from `local_hostname` by removing the port. For e A client is configured through one of the **methods** introduced in [Start a Store Client](#start-a-store-client), plus a shared family of engine-tuning variables: -- **Method A — Programmatic (`setup()` arguments)**: you pass configuration as explicit Python arguments. `MOONCAKE_*` variables are **not** read in this method. +- **Method A — Programmatic (`setup()` arguments)**: launcher-level fields are passed as explicit Python arguments instead of being loaded through `MooncakeConfig`. Backend-specific variables read by C++, including `MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR` and `MOONCAKE_DFS_*`, still apply. - **Method B — Service / Integration (`MOONCAKE_*` + CLI)**: `mooncake.mooncake_store_service` and the vLLM/SGLang connectors read `MOONCAKE_*` environment variables (via `MooncakeConfig`). - **Method C — Resource-owning real client (`mooncake_client`)**: configured through `mooncake_client` CLI flags (see the **Method C** subsection below). - **Engine runtime tuning (`MC_*`)**: low-level variables read by the C++ Transfer Engine / store client at runtime. They are orthogonal to the above and **apply to all methods**. @@ -756,14 +944,14 @@ Arguments of `MooncakeDistributedStore.setup(...)`: | `rdma_devices` | str | required | RDMA NIC(s), comma-separated (pass `""` for non-RDMA). **Keyword is `rdma_devices`, not `device_name`** | | `master_server_addr` | str | required | Master `host:port`. **Keyword is `master_server_addr`, not `master_server_address`** | | `engine` | TransferEngine | `None` | *(advanced)* Reuse an existing Transfer Engine instance instead of creating one | -| `enable_ssd_offload` | bool | `false` | *(advanced)* Enable client-side SSD offload | -| `ssd_offload_path` | str | empty | *(advanced)* SSD offload directory | +| `enable_ssd_offload` | bool | `false` | *(advanced)* Initialize client-side `FileStorage`; required for SSD offload and descriptor-based DFS | +| `ssd_offload_path` | str | empty | *(advanced)* FileStorage path; with the distributed backend, DFS data uses `MOONCAKE_DFS_ROOT_DIR` | | `tenant_id` | str | `default` | *(advanced)* Tenant identifier | | `enable_client_http_server` | bool | `false` | Enable the client-side HTTP `/health`, `/metrics`, and `/metrics/summary` endpoints | | `client_http_port` | int | `9300` | Client-side HTTP endpoint port, used only when `enable_client_http_server=true` | ```{note} -The first seven arguments have **no Python default** — the C++ defaults are not exposed by the pybind binding, so they must all be supplied (a bare `setup(local_hostname, metadata_server)` raises `TypeError`). The later arguments (`engine`, SSD offload fields, `tenant_id`, and client HTTP endpoint fields) are optional. Also, in Method A the `MOONCAKE_*` variables used by `MooncakeConfig` are ignored; low-level runtime variables such as the `MC_*` engine variables below are still read by the C++ client. +The first seven arguments have **no Python default** — the C++ defaults are not exposed by the pybind binding, so they must all be supplied (a bare `setup(local_hostname, metadata_server)` raises `TypeError`). The later arguments (`engine`, SSD offload fields, `tenant_id`, and client HTTP endpoint fields) are optional. In Method A, launcher-level `MOONCAKE_*` variables used only by `MooncakeConfig` are ignored. Variables consumed directly by the C++ client, including the FileStorage/DFS backend variables and low-level `MC_*` engine variables below, are still read. ``` ### Method B — Service / Integration (`MOONCAKE_*` + CLI) @@ -787,8 +975,8 @@ The store service CLI only accepts `--config`, `-D/--define`, `--port`, and `--m | `MOONCAKE_GLOBAL_SEGMENT_SIZE` | `global_segment_size` | `3355443200` (3.125 GiB) | DRAM contributed; accepts byte integer **or** suffixed form like `500gb` | | `MOONCAKE_LOCAL_BUFFER_SIZE` | `local_buffer_size` | `1073741824` (1 GiB) | Transfer Engine buffer; same parsing as above | | `MOONCAKE_LOCAL_HOSTNAME` | `local_hostname` | `localhost` | | -| `MOONCAKE_OFFLOAD_ENABLED` | `enable_ssd_offload` | `false` | Client-side SSD offload | -| `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` | `ssd_offload_path` | empty | Offload directory | +| `MOONCAKE_OFFLOAD_ENABLED` | `enable_ssd_offload` | `false` | Initialize client-side `FileStorage`; required for SSD offload and descriptor-based DFS | +| `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` | `ssd_offload_path` | empty | FileStorage path; DFS shard data uses `MOONCAKE_DFS_ROOT_DIR` with the distributed backend | | `MOONCAKE_TENANT_ID` | `tenant_id` | `default` | Tenant identifier | | `MOONCAKE_ENABLE_CLIENT_HTTP_SERVER` | `enable_client_http_server` | `false` | Enable client-side `/health`, `/metrics`, and `/metrics/summary` endpoints | | `MOONCAKE_CLIENT_HTTP_PORT` | `client_http_port` | `9300` | Client-side HTTP endpoint port | diff --git a/docs/source/design/mooncake-store.md b/docs/source/design/mooncake-store.md index 557d40ed88..98b6f1138a 100644 --- a/docs/source/design/mooncake-store.md +++ b/docs/source/design/mooncake-store.md @@ -685,15 +685,19 @@ Mooncake Store provides a **preferred segment allocation** feature that allows u ### How It Works -The preferred segment allocation feature is implemented through the `AllocationStrategy` system and is controlled via the `preferred_segment` field in the `ReplicateConfig` structure: +The preferred segment allocation feature is implemented through the +`AllocationStrategy` system. The following excerpt shows the legacy +single-segment field and the memory-replica count used by this path; other +`ReplicateConfig` fields are omitted: ```cpp struct ReplicateConfig { - size_t replica_num{1}; // Total number of replicas for the object + size_t replica_num{1}; // Number of memory replicas SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; std::optional soft_pin_ttl_ms{}; // ENABLE override; omitted uses the Master default bool with_hard_pin{false}; // Whether to enable hard pin (never evicted) std::string preferred_segment{}; // Preferred segment for allocation + // Other fields, including nof_replica_num and dfs_replica_num, are omitted. }; ``` @@ -710,38 +714,80 @@ When a `Put` operation is initiated with a non-empty `preferred_segment` value, ## Multi-layer Storage Support -This system provides support for a hierarchical cache architecture, enabling efficient data access through a combination of in-memory caching and persistent storage. Data is initially stored in memory cache and asynchronously backed up to a Distributed File System (DFS), forming a two-tier "memory-SSD persistent storage" cache structure. +Mooncake Store supports three file-backed storage models in addition to memory +and NoF replicas: -### Enabling Persistence Functionality +- `DISK` is the legacy shared-filesystem persistence path enabled by the + master's `--root_fs_dir` flag. +- `LOCAL_DISK` replicas are owned by a real client and use the asynchronous SSD + offload and heartbeat protocol. +- `DFS` replicas occupy globally allocated ranges in shard files on a shared + distributed filesystem. Any correctly configured client can access them. -When the user specifies `--root_fs_dir=/path/to/dir` when starting the master, and this path is a valid DFS-mounted directory on all machines where the clients reside, Mooncake Store's tiered caching functionality will work properly. Additionally, during master initialization, a `cluster_id` is loaded. This ID can be specified during master initialization (`--cluster_id=xxxx`). If not specified, the default value `mooncake_cluster` will be used. Subsequently, the root directory for client persistence will be `/`. +These models have independent placement, metadata, and lifecycle rules. -​Note​​: When enabling this feature, the user must ensure that the DFS-mounted directory (`root_fs_dir=/path/to/dir`) is valid and consistent across all client hosts. If some clients have invalid or incorrect mount paths, it may cause abnormal behavior in Mooncake Store. +### Legacy shared-filesystem `DISK` replicas -This `root_fs_dir` path is a legacy persistence path. SSD offload uses `--enable_offload=true` on the master and real client, stores data under the real client's `MOONCAKE_OFFLOAD_FILE_STORAGE_PATH`, and records `LOCAL_DISK` replicas. Do not use `--root_fs_dir` with `--enable_offload=true`. +When the master starts with `--root_fs_dir=/shared/path`, it adds a legacy +`DISK` replica to each new object. Clients write that replica asynchronously to +the per-cluster directory `/` and can read it when the +normal replica-selection path chooses `DISK`. The path must identify the same +shared filesystem location on every participating client. -### Persistent Storage Space Configuration​ -Mooncake provides configurable DFS available space. Users can specify `--global_file_segment_size=1048576` when starting the master, indicating a maximum usable space of 1MB on DFS. -The current default setting is the maximum value of int64 (as we generally do not restrict DFS storage usage), which is displayed as `infinite` in `mooncake_maseter`'s console logs. -**Notice** The DFS cache space configuration must be used together with the `--root_fs_dir` parameter. Otherwise, you will observe that the `SSD Storage` usage consistently shows: `0 B / 0 B` -**Notice** The capability for file eviction on DFS has not been provided yet +`--global_file_segment_size` declares the legacy file capacity used by master +metrics; its default is unlimited. It does not configure or limit the +descriptor-based DFS shard allocator. Do not combine the legacy path with the +client-owned SSD-offload mode; see the deployment guide for the corresponding +flags and restrictions. -### Data Access Mechanism +### Descriptor-based DFS replicas -The persistence feature also follows Mooncake Store's design principle of separating control flow from data flow. The read/write operations of kvcache objects are completed on the client side, while the query and management functions of kvcache objects are handled on the master side. In the file system, the key -> kvcache object index information is maintained by a fixed indexing mechanism, with each file corresponding to one kvcache object (the filename serves as the associated key name). - -After enabling the persistence feature: - -- For each `Put` or `BatchPut` operation, both a synchronous memory pool write operation and an asynchronous DFS persistence operation will be initiated. -- For each `Get` or `BatchGet` operation, if the corresponding kvcache is not found in the memory pool, the system will attempt to read the file data from DFS and return it to the user. +```{warning} +**Work in progress.** Descriptor-based DFS is not production-ready and is not +covered by the general fault-tolerance, HA continuity, durability, or +multi-tenant guarantees described elsewhere in this design document. +``` -### 3FS USRBIO Plugin (Experimental) +The master owns DFS placement metadata and an allocator for the shared shard +files. During `PutStart` or `UpsertStart`, it allocates an aligned range and +returns a descriptor containing the shard path, shard index, offset, object +size, and aligned size. The replica remains `PROCESSING` until the request is +finalized. Removal, revocation, replacement, and allocator eviction release +the range, with a configurable deferred-free interval preventing immediate +offset reuse. + +The client owns the DFS data plane. `DistributedStorageBackend` validates the +descriptor and delegates positional I/O to either `PosixFsAdapter` or +`Hf3fsAdapter`. The master and clients must use the same DFS root and shard +layout so that a descriptor identifies the same physical file everywhere. + +For a write, the client first completes the requested memory and NoF transfers, +then writes the DFS replica. `Put`, `BatchPut`, `Upsert`, and `BatchUpsert` +acknowledge success only after the requested DFS `WriteAt` operations return +successfully. This is request-synchronous acknowledgement, not an `fsync` +durability guarantee. If either an existing object or an incoming same-size +`Upsert` has a DFS replica, the requested memory, NoF, and DFS replica counts +must match the existing topology. + +For a read, the master returns the readable replica list through the normal +query path. The client selects the first complete replica; if that replica is a +DFS replica, it uses the descriptor to read the requested range directly from +the shared shard file. + +See the {ref}`Mooncake Store deployment guide ` for +configuration, usage, and current limitations. + +### HF3FS USRBIO Adapter (Experimental) ```{note} This integration is **experimental** and incomplete; see the plugin page for details before relying on it. ``` -If you need to use 3FS's native API (USRBIO) to achieve high-performance persistent file reads and writes, you can refer to the configuration instructions in this document [3FS USRBIO Plugin](../getting_started/plugin-usage/3FS-USRBIO-Plugin.md). +The descriptor-based DFS data plane can use the native HF3FS USRBIO API instead +of POSIX I/O. Select it with `MOONCAKE_DFS_FS_ADAPTER=hf3fs`; the legacy +`--root_fs_dir` option does not enable this path, and there is no automatic +fallback to POSIX. See the [HF3FS USRBIO adapter guide](../getting_started/plugin-usage/3FS-USRBIO-Plugin.md) +for build prerequisites and configuration. ## Builtin Metadata Server Mooncake Store provides a built-in HTTP metadata server as an alternative to etcd for storing cluster metadata. This feature is particularly useful for development environments or scenarios where etcd is not available. diff --git a/docs/source/getting_started/plugin-usage/3FS-USRBIO-Plugin.md b/docs/source/getting_started/plugin-usage/3FS-USRBIO-Plugin.md index 5566bc274b..d933af3782 100644 --- a/docs/source/getting_started/plugin-usage/3FS-USRBIO-Plugin.md +++ b/docs/source/getting_started/plugin-usage/3FS-USRBIO-Plugin.md @@ -1,47 +1,92 @@ -# Mooncake HF3FS Plugin (Experimental) +# Mooncake HF3FS USRBIO Adapter (Experimental) ```{warning} -**Experimental / incomplete.** The HF3FS (3FS USRBIO) integration is under development and is not yet considered production-ready. Behavior, build flags, and configuration may change without notice. Use only for evaluation and testing. +**Work in progress / experimental.** Descriptor-based DFS and its HF3FS (3FS +USRBIO) adapter are under development and are not production-ready. They are +not covered by Mooncake Store's general fault-tolerance, HA continuity, +durability, or multi-tenant guarantees. Behavior, build flags, and +configuration may change without notice. Use only for evaluation and testing. ``` -This plugin implements 3FS native API (USRBIO) as a high-performance storage backend for Mooncake. +This adapter implements the HF3FS native USRBIO data plane for Mooncake Store's +descriptor-based DFS replicas. The master allocates ranges in shared shard +files, and clients use USRBIO to access the ranges described by the replica +metadata. + +The adapter is not enabled by the legacy `--root_fs_dir` option. It also does +not automatically fall back to POSIX I/O; select +`MOONCAKE_DFS_FS_ADAPTER=posix` explicitly when POSIX behavior is required. ## Prerequisites -### 1. 3FS Installation +### 1. HF3FS installation + - Build and install [3FS](https://github.com/deepseek-ai/3FS/) - Required library: `libhf3fs_api_shared.so` (Default location: `3FS_PATH/build/src/lib/api`) → Install to: `/usr/lib/` - Required header: `hf3fs_usrbio.h` (Default location: `3FS_PATH/src/lib/api`) → Install to: `/usr/include/` -### 2. Mooncake Configuration -- Enable 3FS support during CMake configuration: -```bash +### 2. Mooncake build + +Enable HF3FS support during CMake configuration: +```bash cmake -DUSE_3FS=ON ... ``` -- Build and install Mooncake as usual. +Then build and install Mooncake as usual. ## Usage -### Basic Operation -Start master server and specify the 3FS mount point: +### Master + +Enable descriptor-based DFS and point it at a directory on the shared HF3FS +mount: + ```bash +export MOONCAKE_ENABLE_DFS=1 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true -./build/mooncake-store/src/mooncake_master \ - --root_fs_dir=/path/to/3fs_mount_point +./build/mooncake-store/src/mooncake_master [other master arguments] ``` -### Important Notes -1. The specified directory **must** be a 3FS mount point - - If not, the system will automatically fall back to POSIX API -2. For optimal performance: - - Ensure proper permissions on the 3FS mount point - - Verify 3FS service is running before execution - -### Example + +The master creates and preallocates the configured shard files during startup. +Ensure the mount is available, writable, and has enough capacity before +starting the process. + +### Clients + +Every client that may read or write DFS replicas must use the same root, +adapter, shard count, shard capacity, and alignment. The DFS root must be an +absolute path and use the same path string in every process. For the standalone +store service: + ```bash +export MOONCAKE_OFFLOAD_ENABLED=true +export MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR=distributed_storage_backend +export MOONCAKE_OFFLOAD_FILE_STORAGE_PATH=/data/file_storage +export MOONCAKE_MASTER=127.0.0.1:50051 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true -ROLE=prefill MOONCAKE_STORAGE_ROOT_DIR=/mnt/3fs python3 ./stress_cluster_benchmark.py +python -m mooncake.mooncake_store_service ``` + +`MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` must already be an absolute, writable, +non-symlink directory. DFS shard data is stored under +`MOONCAKE_DFS_ROOT_DIR`; the separate FileStorage path is still validated +during client initialization. + +For the complete configuration reference, request example, synchronous write +semantics, and current recovery limitations, see the {ref}`DFS deployment +documentation `. diff --git a/mooncake-common/include/environ.h b/mooncake-common/include/environ.h index f6224696dd..18fcf493a1 100644 --- a/mooncake-common/include/environ.h +++ b/mooncake-common/include/environ.h @@ -1,8 +1,8 @@ #pragma once -#include #include #include +#include namespace mooncake { @@ -96,6 +96,7 @@ class Environ { static int64_t GetInt64(const char* name, int64_t default_value); static uint32_t GetUInt32(const char* name, uint32_t default_value); static uint64_t GetUInt64(const char* name, uint64_t default_value); + static double GetDouble(const char* name, double default_value); // Helper method to get size_t from env static size_t GetSizeT(const char* name, size_t default_value); // Helper method to get a canonical boolean from env. Invalid values use the diff --git a/mooncake-common/src/environ.cpp b/mooncake-common/src/environ.cpp index c736823171..31b37e21c4 100644 --- a/mooncake-common/src/environ.cpp +++ b/mooncake-common/src/environ.cpp @@ -1,6 +1,9 @@ #include "environ.h" #include +#include +#include +#include #include #include #include @@ -64,6 +67,29 @@ size_t ReadSizeT(const EnvironSource& source, const char* name, return ReadInteger(source, name, default_value); } +double ReadDouble(const EnvironSource& source, const char* name, + double default_value) { + const char* value = source.Get(name); + if (value == nullptr || value[0] == '\0') { + return default_value; + } + + char* end = nullptr; + errno = 0; + const double parsed = std::strtod(value, &end); + while (end != nullptr && std::isspace(static_cast(*end))) { + ++end; + } + if (end != value && end != nullptr && *end == '\0' && errno != ERANGE && + std::isfinite(parsed)) { + return parsed; + } + + std::cerr << "[Mooncake] Warning: invalid value '" << value << "' for env " + << name << ", using default " << default_value << std::endl; + return default_value; +} + bool ReadBool(const EnvironSource& source, const char* name, bool default_value) { const char* value = source.Get(name); @@ -117,6 +143,10 @@ uint64_t Environ::GetUInt64(const char* name, uint64_t default_value) { return ReadInteger(GetOsEnvironSource(), name, default_value); } +double Environ::GetDouble(const char* name, double default_value) { + return ReadDouble(GetOsEnvironSource(), name, default_value); +} + size_t Environ::GetSizeT(const char* name, size_t default_value) { return ReadSizeT(GetOsEnvironSource(), name, default_value); } diff --git a/mooncake-common/tests/environ_test.cpp b/mooncake-common/tests/environ_test.cpp index 4701f6e78e..c1c5812959 100644 --- a/mooncake-common/tests/environ_test.cpp +++ b/mooncake-common/tests/environ_test.cpp @@ -32,6 +32,7 @@ class EnvironTest : public ::testing::Test { unsetenv("MC_TEST_UINT32"); unsetenv("MC_TEST_UINT64"); unsetenv("MC_TEST_SIZET"); + unsetenv("MC_TEST_DOUBLE"); unsetenv("MC_TEST_BOOL"); unsetenv("MC_TEST_STRING"); // Make sure AWS vars don't leak in from the test runner's env. @@ -139,6 +140,21 @@ TEST_F(EnvironTest, UnsignedGettersUseRequestedDefaultForInvalidValues) { EXPECT_EQ(Environ::GetUInt64("MC_TEST_UINT64", 23), 23U); } +// --- GetDouble --- + +TEST_F(EnvironTest, GetDoubleValidValue) { + setenv("MC_TEST_DOUBLE", " 0.75 ", 1); + EXPECT_DOUBLE_EQ(Environ::GetDouble("MC_TEST_DOUBLE", 0.5), 0.75); +} + +TEST_F(EnvironTest, GetDoubleMissingOrInvalidUsesRequestedDefault) { + EXPECT_DOUBLE_EQ(Environ::GetDouble("MC_TEST_DOUBLE", 0.5), 0.5); + setenv("MC_TEST_DOUBLE", "0.75garbage", 1); + EXPECT_DOUBLE_EQ(Environ::GetDouble("MC_TEST_DOUBLE", 0.5), 0.5); + setenv("MC_TEST_DOUBLE", "nan", 1); + EXPECT_DOUBLE_EQ(Environ::GetDouble("MC_TEST_DOUBLE", 0.5), 0.5); +} + // --- AWS / S3 fields --- // // NOTE: Environ is a singleton whose constructor caches every value the diff --git a/mooncake-integration/store/store_py.cpp b/mooncake-integration/store/store_py.cpp index fdeb35432d..3f01689a39 100644 --- a/mooncake-integration/store/store_py.cpp +++ b/mooncake-integration/store/store_py.cpp @@ -1924,6 +1924,7 @@ PYBIND11_MODULE(store, m) { .def(py::init<>()) .def_readwrite("replica_num", &ReplicateConfig::replica_num) .def_readwrite("nof_replica_num", &ReplicateConfig::nof_replica_num) + .def_readwrite("dfs_replica_num", &ReplicateConfig::dfs_replica_num) .def_readwrite("soft_pin_action", &ReplicateConfig::soft_pin_action) .def_readwrite("soft_pin_ttl_ms", &ReplicateConfig::soft_pin_ttl_ms) .def_readwrite("with_hard_pin", &ReplicateConfig::with_hard_pin) diff --git a/mooncake-integration/store/store_py_internal.h b/mooncake-integration/store/store_py_internal.h index e7f439f443..e43c3df105 100644 --- a/mooncake-integration/store/store_py_internal.h +++ b/mooncake-integration/store/store_py_internal.h @@ -861,7 +861,8 @@ bool parallelism_specs_equal_by_kind(const TensorParallelismSpec &lhs, } bool is_default_replicate_config(const ReplicateConfig &config) { - return config.replica_num == 1 && + return config.replica_num == 1 && config.nof_replica_num == 0 && + config.dfs_replica_num == 0 && config.soft_pin_action == SoftPinAction::PRESERVE && !config.soft_pin_ttl_ms.has_value() && !config.with_hard_pin && config.preferred_segments.empty() && diff --git a/mooncake-store/include/allocator.h b/mooncake-store/include/allocator.h index e0321f5c6f..5486664a8c 100644 --- a/mooncake-store/include/allocator.h +++ b/mooncake-store/include/allocator.h @@ -20,11 +20,12 @@ namespace mooncake { * @brief Type of buffer allocator used in the system */ enum class ReplicaType { - MEMORY, // Memory replica - DISK, // Disk replica - LOCAL_DISK, // Local disk replica - NOF_SSD, // Nvme-oF SSD replica - ALL, // All memory and NoF replicas in put finalize path + MEMORY = 0, // Memory replica + DISK = 1, // Disk replica + LOCAL_DISK = 2, // Local disk replica + NOF_SSD = 3, // Nvme-oF SSD replica + ALL = 4, // All synchronous replicas in put finalize path + DFS = 100, // Distributed filesystem page-offset replica }; // Constant for unknown free space in allocators that don't track it precisely diff --git a/mooncake-store/include/client_service.h b/mooncake-store/include/client_service.h index 094f8d8cc6..03522dcc38 100644 --- a/mooncake-store/include/client_service.h +++ b/mooncake-store/include/client_service.h @@ -31,6 +31,7 @@ namespace mooncake { class PutOperation; +class DistributedStorageBackend; class RealClient; /** @@ -556,6 +557,8 @@ class Client { tl::expected NotifyOffloadSuccess( const std::vector& tasks, const std::vector& metadatas); + void SetDfsStorageBackend( + std::shared_ptr backend); /** * @brief Fetch tasks assigned to a client @@ -775,6 +778,9 @@ class Client { ErrorCode TransferReadRange(const Replica::Descriptor& replica_descriptor, std::vector& slices, uint64_t src_offset); + ErrorCode ReadDfsReplica(const std::string& key, + const Replica::Descriptor& replica_descriptor, + std::vector& slices); tl::expected ComputeObjectChecksumForSlices( const std::string& object_key, const std::vector& slices, size_t object_size); @@ -852,6 +858,7 @@ class Client { void ComputeBatchObjectChecksums(std::vector& ops); void SubmitTransfers(std::vector& ops); void WaitForTransfers(std::vector& ops); + void SubmitDfsWrites(std::vector& ops); void FinalizeBatchPut(std::vector& ops); void StartBatchUpsert(std::vector& ops, const ReplicateConfig& config); @@ -859,6 +866,11 @@ class Client { std::vector> CollectResults( const std::vector& ops); + std::vector WriteDfsReplicas( + const std::vector& keys, + const std::vector*>& slice_lists, + const std::vector& descriptors); + std::vector> BatchPutWhenPreferSameNode( std::vector& ops); std::vector> BatchGetWhenPreferSameNode( @@ -918,6 +930,7 @@ class Client { std::unique_ptr pinned_buffer_pool_; ThreadPool write_thread_pool_; std::shared_ptr storage_backend_; + std::shared_ptr dfs_storage_backend_; // For high availability std::unique_ptr leader_coordinator_; diff --git a/mooncake-store/include/hf3fs/hf3fs.h b/mooncake-store/include/hf3fs/hf3fs.h index b624d15623..feebaab516 100644 --- a/mooncake-store/include/hf3fs/hf3fs.h +++ b/mooncake-store/include/hf3fs/hf3fs.h @@ -5,12 +5,12 @@ #include #include #include + +#include "file_interface.h" #include "types.h" namespace mooncake { -class StorageFile; - // Forward declaration of USRBIOResourceManager struct Hf3fsConfig { // 3FS cluster related parameters @@ -103,4 +103,4 @@ class ThreeFSFile : public StorageFile { USRBIOResourceManager *resource_manager_; }; -} // namespace mooncake \ No newline at end of file +} // namespace mooncake diff --git a/mooncake-store/include/master_service.h b/mooncake-store/include/master_service.h index a58048f83d..af38f4662a 100644 --- a/mooncake-store/include/master_service.h +++ b/mooncake-store/include/master_service.h @@ -61,6 +61,9 @@ struct MasterSnapshotPayloads; class MasterSnapshotCodecTest; // test fixture, needs private state access } // namespace ha +class EtcdOpLogStore; +class DfsGlobalAllocator; + // Forward declarations class AllocationStrategy; class EvictionStrategy; @@ -172,6 +175,7 @@ class MasterService { double evict_ratio_lowerbound); void RunNoFBatchEvictForTesting(double evict_ratio_target, double evict_ratio_lowerbound); + void RunDfsEvictionForTesting(); /** * @brief Mount a memory segment for buffer allocation. This function is @@ -1154,7 +1158,7 @@ class MasterService { return EraseReplicas([replica_type](const Replica& replica) { if (replica_type == ReplicaType::ALL) { return replica.is_memory_replica() || - replica.is_nof_replica(); + replica.is_nof_replica() || replica.is_dfs_replica(); } return replica.type() == replica_type; }); @@ -1808,7 +1812,10 @@ class MasterService { void DiscardExpiredProcessingReplicas( MetadataShardAccessorRW& shard, const std::chrono::system_clock::time_point& now); - + void FreeDfsReplicas(const std::string& key, + const std::vector& replicas); + void RunDfsEviction(); + void InitDfsAllocatorFromEnvironment(const MasterServiceConfig& config); /** * @brief Helper to release space of expired discarded replicas. * @return Number of released objects that have memory replicas @@ -2363,6 +2370,8 @@ class MasterService { void cleanupHttpMetadata(const std::string& segment_name); bool use_disk_replica_{false}; + bool enable_dfs_{false}; + std::unique_ptr dfs_allocator_; // Segment management SegmentManager segment_manager_; diff --git a/mooncake-store/include/replica.h b/mooncake-store/include/replica.h index 49696e5e50..95c8072b4e 100644 --- a/mooncake-store/include/replica.h +++ b/mooncake-store/include/replica.h @@ -37,7 +37,8 @@ inline std::ostream& operator<<(std::ostream& os, {ReplicaType::DISK, "DISK"}, {ReplicaType::LOCAL_DISK, "LOCAL_DISK"}, {ReplicaType::NOF_SSD, "NOF_SSD"}, - {ReplicaType::ALL, "ALL"}}; + {ReplicaType::ALL, "ALL"}, + {ReplicaType::DFS, "DFS"}}; os << (replica_type_strings.count(replicaType) ? replica_type_strings.at(replicaType) @@ -103,6 +104,7 @@ inline std::ostream& operator<<(std::ostream& os, struct ReplicateConfig { size_t replica_num{1}; size_t nof_replica_num{0}; + size_t dfs_replica_num{0}; SoftPinAction soft_pin_action{SoftPinAction::PRESERVE}; // Optional request-level override. When omitted, ENABLE uses the // master's default soft-pin TTL. @@ -135,6 +137,7 @@ struct ReplicateConfig { const ReplicateConfig& config) noexcept { os << "ReplicateConfig: { replica_num: " << config.replica_num << ", nof_replica_num: " << config.nof_replica_num + << ", dfs_replica_num: " << config.dfs_replica_num << ", soft_pin_action: " << config.soft_pin_action << ", soft_pin_ttl_ms: "; if (config.soft_pin_ttl_ms.has_value()) { @@ -186,10 +189,12 @@ enum class ReplicaWriteMode { inline ReplicaWriteMode DetermineReplicaWriteMode( const ReplicateConfig& config) { - if (config.replica_num == 1 && config.nof_replica_num == 1) { + if (config.dfs_replica_num == 0 && config.replica_num == 1 && + config.nof_replica_num == 1) { return ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA; } - if (config.replica_num > 1 || config.nof_replica_num > 1) { + if (config.replica_num > 1 || config.nof_replica_num > 1 || + config.dfs_replica_num > 0) { return ReplicaWriteMode::RELIABLE_MULTI_REPLICA; } return ReplicaWriteMode::SINGLE_REPLICA; @@ -214,6 +219,20 @@ struct LocalDiskReplicaData { std::string transport_endpoint; }; +struct DistributedFSDescriptor { + std::string file_path; + uint64_t offset = 0; + uint64_t object_size = 0; + uint64_t aligned_size = 0; + int shard_idx = 0; + YLT_REFL(DistributedFSDescriptor, file_path, offset, object_size, + aligned_size, shard_idx); +}; + +struct DfsReplicaData { + DistributedFSDescriptor descriptor; +}; + struct MemoryDescriptor { AllocatedBuffer::Descriptor buffer_descriptor; YLT_REFL(MemoryDescriptor, buffer_descriptor); @@ -282,6 +301,13 @@ class Replica { MasterMetricManager::instance().inc_allocated_file_size(object_size); } + // dfs replica constructor + Replica(DistributedFSDescriptor descriptor, ReplicaStatus status) + : id_(next_id_.fetch_add(1)), + data_(DfsReplicaData{std::move(descriptor)}), + status_(status), + refcnt_(0) {} + ~Replica() { if (status_ == ReplicaStatus::UNDEFINED) return; if (is_disk_replica()) { @@ -399,6 +425,22 @@ class Replica { return replica.is_local_disk_replica(); } + [[nodiscard]] bool is_dfs_replica() const { + return std::holds_alternative(data_); + } + + [[nodiscard]] static bool fn_is_dfs_replica(const Replica& replica) { + return replica.is_dfs_replica(); + } + + [[nodiscard]] const DistributedFSDescriptor& get_dfs_descriptor() const { + return std::get(data_).descriptor; + } + + [[nodiscard]] DistributedFSDescriptor& get_dfs_descriptor() { + return std::get(data_).descriptor; + } + [[nodiscard]] bool has_invalid_mem_handle() const { if (is_memory_replica()) { const auto& mem_data = std::get(data_); @@ -516,12 +558,15 @@ class Replica { ReplicaType operator()(const LocalDiskReplicaData&) const { return ReplicaType::LOCAL_DISK; } + ReplicaType operator()(const DfsReplicaData&) const { + return ReplicaType::DFS; + } }; struct Descriptor { ReplicaID id; std::variant + LocalDiskDescriptor, DistributedFSDescriptor> descriptor_variant; ReplicaStatus status; YLT_REFL(Descriptor, id, descriptor_variant, status); @@ -561,6 +606,16 @@ class Replica { descriptor_variant); } + bool is_dfs_replica() noexcept { + return std::holds_alternative( + descriptor_variant); + } + + bool is_dfs_replica() const noexcept { + return std::holds_alternative( + descriptor_variant); + } + MemoryDescriptor& get_memory_descriptor() { if (auto* desc = std::get_if(&descriptor_variant)) { @@ -591,6 +646,14 @@ class Replica { throw std::runtime_error("Expected LocalDiskDescriptor"); } + DistributedFSDescriptor& get_dfs_descriptor() { + if (auto* desc = + std::get_if(&descriptor_variant)) { + return *desc; + } + throw std::runtime_error("Expected DistributedFSDescriptor"); + } + const MemoryDescriptor& get_memory_descriptor() const { if (auto* desc = std::get_if(&descriptor_variant)) { @@ -620,6 +683,14 @@ class Replica { } throw std::runtime_error("Expected LocalDiskDescriptor"); } + + const DistributedFSDescriptor& get_dfs_descriptor() const { + if (auto* desc = + std::get_if(&descriptor_variant)) { + return *desc; + } + throw std::runtime_error("Expected DistributedFSDescriptor"); + } }; private: @@ -627,7 +698,7 @@ class Replica { ReplicaID id_; std::variant + LocalDiskReplicaData, DfsReplicaData> data_; ReplicaStatus status_{ReplicaStatus::UNDEFINED}; @@ -678,6 +749,8 @@ inline Replica::Descriptor Replica::get_descriptor() const { local_disk_desc.object_size = disk_data.object_size; local_disk_desc.transport_endpoint = disk_data.transport_endpoint; desc.descriptor_variant = std::move(local_disk_desc); + } else if (is_dfs_replica()) { + desc.descriptor_variant = std::get(data_).descriptor; } return desc; @@ -729,6 +802,18 @@ inline std::ostream& operator<<(std::ostream& os, const Replica& replica) { const auto& disk_data = std::get(replica.data_); os << "type: DISK, file_path: " << disk_data.file_path << ", object_size: " << disk_data.object_size; + } else if (replica.is_local_disk_replica()) { + const auto& disk_data = std::get(replica.data_); + os << "type: LOCAL_DISK, client_id: " << disk_data.client_id + << ", object_size: " << disk_data.object_size + << ", transport_endpoint: " << disk_data.transport_endpoint; + } else if (replica.is_dfs_replica()) { + const auto& dfs_data = std::get(replica.data_); + os << "type: DFS, file_path: " << dfs_data.descriptor.file_path + << ", offset: " << dfs_data.descriptor.offset + << ", object_size: " << dfs_data.descriptor.object_size + << ", aligned_size: " << dfs_data.descriptor.aligned_size + << ", shard_idx: " << dfs_data.descriptor.shard_idx; } os << ", refcnt: " << replica.refcnt_.load() << " }"; diff --git a/mooncake-store/include/replica_selection.h b/mooncake-store/include/replica_selection.h index ed7a49c2af..eb7fef190b 100644 --- a/mooncake-store/include/replica_selection.h +++ b/mooncake-store/include/replica_selection.h @@ -115,10 +115,11 @@ inline const Replica::Descriptor *PickBestRemoteMemory( return best; } -// Select the best replica from a list: prefer local MEMORY, then any MEMORY, -// then LOCAL_DISK, then DISK. Master may return replicas in any order, so we -// always scan. When scoring is enabled and there are multiple remote MEMORY -// replicas, the best-scoring one is chosen instead of the first encountered. +// Select the best replica from a list: prefer local MEMORY, local NOF_SSD, +// remote MEMORY, remote NOF_SSD, LOCAL_DISK, DFS, then DISK. Master may return +// replicas in any order, so we always scan. When scoring is enabled and there +// are multiple remote MEMORY replicas, the best-scoring one is chosen instead +// of the first encountered. inline const Replica::Descriptor *SelectBestReplica( const std::vector &replicas, const std::unordered_set &local_endpoints) { @@ -157,7 +158,9 @@ inline const Replica::Descriptor *SelectBestReplica( for (const auto &r : replicas) { if (r.status != ReplicaStatus::COMPLETE) continue; if (r.is_local_disk_replica()) { - best = &r; // LOCAL_DISK always overrides DISK + best = &r; // LOCAL_DISK always overrides DFS and DISK + } else if (r.is_dfs_replica()) { + if (!best || !best->is_local_disk_replica()) best = &r; } else if (r.is_disk_replica() && !best) { best = &r; } diff --git a/mooncake-store/include/storage/distributed/dfs_global_allocator.h b/mooncake-store/include/storage/distributed/dfs_global_allocator.h new file mode 100644 index 0000000000..1c8dc540e3 --- /dev/null +++ b/mooncake-store/include/storage/distributed/dfs_global_allocator.h @@ -0,0 +1,166 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "offset_allocator/offset_allocator.h" +#include "replica.h" +#include "storage/distributed/fs_adapter.h" +#include "types.h" + +namespace mooncake { + +struct DistributedStorageConfig; + +class DfsGlobalAllocator { + public: + struct EvictionCandidate { + std::string key; + int shard_idx; + uint64_t offset; + }; + + // Keeps selected allocations pinned while the master decides which + // candidates can be evicted. An unresolved transaction is aborted on + // destruction so candidates cannot get stuck outside the LRU. + class PendingEviction { + public: + PendingEviction() = default; + ~PendingEviction(); + + PendingEviction(const PendingEviction&) = delete; + PendingEviction& operator=(const PendingEviction&) = delete; + PendingEviction(PendingEviction&& other) noexcept; + PendingEviction& operator=(PendingEviction&&) = delete; + + bool Empty() const { return candidates_.empty(); } + const std::vector& Candidates() const { + return candidates_; + } + + private: + friend class DfsGlobalAllocator; + + struct PreparedAllocation { + EvictionCandidate candidate; + std::shared_ptr handle; + uint64_t bytes = 0; + }; + + explicit PendingEviction(DfsGlobalAllocator* owner) : owner_(owner) {} + + DfsGlobalAllocator* owner_ = nullptr; + std::vector candidates_; + std::vector prepared_; + }; + + DfsGlobalAllocator() = default; + ~DfsGlobalAllocator(); + + DfsGlobalAllocator(const DfsGlobalAllocator&) = delete; + DfsGlobalAllocator& operator=(const DfsGlobalAllocator&) = delete; + + tl::expected Init(const DistributedStorageConfig& config); + bool IsInitialized() const { + return initialized_.load(std::memory_order_acquire); + } + + tl::expected Allocate( + const std::string& key, uint64_t size); + void Free(uint64_t offset, uint64_t aligned_size, int shard_idx, + const std::string& key); + void UpdateAccess(const std::string& key, int shard_idx, uint64_t offset); + PendingEviction PrepareEviction(); + void CommitPreparedEviction(PendingEviction&& pending); + void RestorePreparedEviction(PendingEviction&& pending); + void ResolvePreparedEviction(PendingEviction&& pending, + const std::vector& accepted); + + bool IsEvictionEnabled() const { return eviction_enabled_; } + std::chrono::seconds GetEvictionCheckInterval() const { + return eviction_check_interval_; + } + + static std::string FormatShardIdx(int idx, int shard_count); + + private: + using OffsetAllocator = offset_allocator::OffsetAllocator; + using OffsetAllocationHandle = offset_allocator::OffsetAllocationHandle; + + struct ShardState { + uint64_t capacity = 0; + std::shared_ptr allocator; + + struct AllocationRecord { + std::string key; + std::shared_ptr handle; + uint64_t bytes = 0; + bool eviction_prepared = false; + }; + + std::shared_mutex handle_mutex; + std::unordered_map offset_to_handle; + + std::mutex lru_mutex; + std::list> lru_list; + std::unordered_map lru_index; + // Once the high watermark is crossed, keep selecting candidates until + // effective usage falls below the low watermark. Protected candidates + // may make that span multiple prepare/resolve rounds. + bool eviction_active = false; + + struct PendingFree { + std::shared_ptr handle; + uint64_t bytes = 0; + std::chrono::steady_clock::time_point when; + }; + std::mutex pending_mutex; + std::deque pending_free; + uint64_t pending_free_bytes = 0; + }; + + static constexpr size_t kNumKeyStripes = 65536; + + std::unique_lock LockKey(const std::string& key) { + return std::unique_lock( + key_stripes_[std::hash{}(key) % kNumKeyStripes]); + } + + void ProcessPendingFrees(int shard_idx); + void QueuePendingFree(ShardState& shard, + const std::shared_ptr& handle, + uint64_t bytes, + std::chrono::steady_clock::time_point when); + void CleanupExpiredPendingFrees(ShardState& shard, + std::chrono::steady_clock::time_point now); + double EffectiveUsage(ShardState& shard); + void PrepareEvictionFromShard(int shard_idx, PendingEviction& pending); + int SelectShard(const std::string& key) const; + uint64_t AlignSize(uint64_t size) const; + + std::string mount_path_; + int shard_count_ = 0; + uint64_t alignment_ = 4096; + std::vector> shards_; + std::unique_ptr fs_adapter_; + bool eviction_enabled_ = true; + double eviction_high_watermark_ = 0.9; + double eviction_low_watermark_ = 0.7; + std::chrono::seconds deferred_free_duration_{30}; + std::chrono::seconds eviction_check_interval_{5}; + std::atomic initialized_{false}; + std::array key_stripes_; +}; + +} // namespace mooncake diff --git a/mooncake-store/include/storage/distributed/distributed_storage_backend.h b/mooncake-store/include/storage/distributed/distributed_storage_backend.h index fe733b7fde..ca3aaf27cb 100644 --- a/mooncake-store/include/storage/distributed/distributed_storage_backend.h +++ b/mooncake-store/include/storage/distributed/distributed_storage_backend.h @@ -1,20 +1,48 @@ #pragma once +#include +#include #include +#include +#include +#include #include "fs_adapter.h" +#include "replica.h" #include "storage_backend.h" namespace mooncake { struct DistributedStorageConfig { - std::string fsdir = "distributed_dir"; + std::string fsdir = "/mnt/3fs/mooncake"; std::string fs_adapter_type = "hf3fs"; bool enable_health_check = false; - int hash_bucket_count = 256; + int shard_count = 64; + uint64_t shard_capacity = 4ULL * 1024 * 1024 * 1024; + uint64_t alignment = 4096; + bool single_tenant = true; + bool eviction_enabled = true; + double eviction_high_watermark = 0.9; + double eviction_low_watermark = 0.7; + std::chrono::seconds deferred_free_duration{30}; + std::chrono::seconds eviction_check_interval{5}; bool Validate() const; + bool ValidateForAllocator() const; static DistributedStorageConfig FromEnvironment(); + std::string FormatStr() const; +}; + +struct DfsWriteRequest { + std::string key; + DistributedFSDescriptor descriptor; + std::vector slices; +}; + +struct DfsReadRequest { + std::string key; + DistributedFSDescriptor descriptor; + std::vector slices; }; /** @@ -29,6 +57,7 @@ class DistributedStorageBackend : public StorageBackendInterface { const FileStorageConfig& file_storage_config, const DistributedStorageConfig& distributed_config, std::unique_ptr fs_adapter); + ~DistributedStorageBackend() override; tl::expected Init() override; @@ -39,6 +68,14 @@ class DistributedStorageBackend : public StorageBackendInterface { complete_handler, EvictionHandler eviction_handler = nullptr) override; + std::vector> BatchWrite( + const std::vector& requests); + + std::vector> BatchRead( + const std::vector& requests); + + // Key-only storage backend operations cannot safely address DFS objects; + // callers must use BatchRead/BatchWrite with request-scoped descriptors. tl::expected BatchLoad( std::unordered_map& batched_slices) override; @@ -52,14 +89,16 @@ class DistributedStorageBackend : public StorageBackendInterface { std::vector& metadatas)>& handler) override; private: - std::string GetObjectPath(const std::string& key) const; - static std::string EscapeFilename(const std::string& key); - static std::string UnescapeFilename(const std::string& name); + struct ShardFile { + std::string path; + int fd = -1; + std::mutex mutex; + }; std::unique_ptr fs_adapter_; DistributedStorageConfig distributed_config_; std::string root_dir_; - int hash_bucket_count_; + std::vector> shard_files_; bool initialized_ = false; }; diff --git a/mooncake-store/include/storage/distributed/fs_adapter.h b/mooncake-store/include/storage/distributed/fs_adapter.h index defda61647..c55ce609bf 100644 --- a/mooncake-store/include/storage/distributed/fs_adapter.h +++ b/mooncake-store/include/storage/distributed/fs_adapter.h @@ -97,6 +97,34 @@ class FileSystemAdapter { return result; } + // === fd-based I/O for shard-offset DFS mode === + + virtual tl::expected OpenFile(const std::string& /*path*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); + } + + virtual tl::expected CloseFile(int /*fd*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); + } + + virtual tl::expected PreallocateFile( + const std::string& /*path*/, uint64_t /*size*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); + } + + virtual tl::expected WriteAt(int /*fd*/, + const iovec* /*iov*/, + int /*iovcnt*/, + int64_t /*offset*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); + } + + virtual tl::expected ReadAt(int /*fd*/, iovec* /*iov*/, + int /*iovcnt*/, + int64_t /*offset*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); + } + // === Lifecycle === virtual tl::expected Init( diff --git a/mooncake-store/include/storage/distributed/hf3fs_adapter.h b/mooncake-store/include/storage/distributed/hf3fs_adapter.h index 4b0fa05f8f..a36e6c719b 100644 --- a/mooncake-store/include/storage/distributed/hf3fs_adapter.h +++ b/mooncake-store/include/storage/distributed/hf3fs_adapter.h @@ -38,6 +38,20 @@ class Hf3fsAdapter : public FileSystemAdapter { tl::expected, ErrorCode> ListFiles( const std::string& dir) override; + tl::expected OpenFile(const std::string& path) override; + + tl::expected CloseFile(int fd) override; + + tl::expected PreallocateFile(const std::string& path, + uint64_t size) override; + + tl::expected WriteAt(int fd, const iovec* iov, + int iovcnt, + int64_t offset) override; + + tl::expected ReadAt(int fd, iovec* iov, int iovcnt, + int64_t offset) override; + tl::expected Init(const std::string& mount_path) override; tl::expected Shutdown() override; diff --git a/mooncake-store/include/storage/distributed/posix_fs_adapter.h b/mooncake-store/include/storage/distributed/posix_fs_adapter.h new file mode 100644 index 0000000000..01a3be1f04 --- /dev/null +++ b/mooncake-store/include/storage/distributed/posix_fs_adapter.h @@ -0,0 +1,55 @@ +#pragma once + +#include "storage/distributed/fs_adapter.h" + +namespace mooncake { + +class PosixFsAdapter : public FileSystemAdapter { + public: + tl::expected WriteFile( + const std::string& path, std::span data) override; + + tl::expected ReadFile(const std::string& path, void* buf, + size_t len) override; + + tl::expected VectorWriteFile(const std::string& path, + const iovec* iov, + int iovcnt, + off_t offset) override; + + tl::expected VectorReadFile(const std::string& path, + const iovec* iov, int iovcnt, + off_t offset) override; + + tl::expected DeleteFile(const std::string& path) override; + + tl::expected FileExists(const std::string& path) override; + + tl::expected, ErrorCode> ListFiles( + const std::string& dir) override; + + tl::expected OpenFile(const std::string& path) override; + + tl::expected CloseFile(int fd) override; + + tl::expected PreallocateFile(const std::string& path, + uint64_t size) override; + + tl::expected WriteAt(int fd, const iovec* iov, + int iovcnt, + int64_t offset) override; + + tl::expected ReadAt(int fd, iovec* iov, int iovcnt, + int64_t offset) override; + + tl::expected Init(const std::string& mount_path) override; + + tl::expected Shutdown() override; + + const char* GetName() const override { return "posix"; } + + private: + std::string mount_path_; +}; + +} // namespace mooncake diff --git a/mooncake-store/include/storage_backend.h b/mooncake-store/include/storage_backend.h index b6acb75a6a..c8b61c41b4 100644 --- a/mooncake-store/include/storage_backend.h +++ b/mooncake-store/include/storage_backend.h @@ -335,6 +335,8 @@ struct FileStorageConfig { // Use io_uring for file I/O instead of POSIX pread/pwrite bool use_uring = false; + // DFS page-offset mode. Enabled automatically for kDistributed. + bool enable_dfs = false; // Proactively evict local disk objects from the heartbeat thread once // backend usage crosses the high watermark. bool enable_disk_watermark_eviction = true; @@ -361,6 +363,7 @@ struct FileStorageConfig { class StorageBackendInterface { public: StorageBackendInterface(const FileStorageConfig& file_storage_config); + virtual ~StorageBackendInterface() = default; using EvictionHandler = std::function( const std::vector& evicted_keys)>; diff --git a/mooncake-store/include/types.h b/mooncake-store/include/types.h index 0721098bbb..0793522375 100644 --- a/mooncake-store/include/types.h +++ b/mooncake-store/include/types.h @@ -361,7 +361,8 @@ enum class ErrorCode : int32_t { UNAVAILABLE_IN_CURRENT_STATUS = -1010, ///< Request cannot be done in current status. UNAVAILABLE_IN_CURRENT_MODE = - -1011, ///< Request cannot be done in current mode. + -1011, ///< Request cannot be done in current mode. + NOT_SUPPORTED = -1012, ///< Operation is not supported in current mode. // FILE errors (Range: -1100 to -1199) FILE_NOT_FOUND = -1100, ///< File not found. diff --git a/mooncake-store/include/utils.h b/mooncake-store/include/utils.h index 41d6e3d205..9943f71595 100644 --- a/mooncake-store/include/utils.h +++ b/mooncake-store/include/utils.h @@ -9,6 +9,7 @@ #include #include #include +#include #include #include "rpc_types.h" diff --git a/mooncake-store/src/CMakeLists.txt b/mooncake-store/src/CMakeLists.txt index b747a37cb9..7cbed43f9c 100644 --- a/mooncake-store/src/CMakeLists.txt +++ b/mooncake-store/src/CMakeLists.txt @@ -35,7 +35,9 @@ set(MOONCAKE_STORE_SOURCES http_metadata_server.cpp file_storage.cpp serialize/serializer.cpp + storage/distributed/dfs_global_allocator.cpp storage/distributed/distributed_storage_backend.cpp + storage/distributed/posix_fs_adapter.cpp device/accelerator_device.cpp device/accelerator_registry.cpp device/runtime_accelerator.cpp diff --git a/mooncake-store/src/client_buffer.cpp b/mooncake-store/src/client_buffer.cpp index 7ac1dc9212..9206a93c46 100644 --- a/mooncake-store/src/client_buffer.cpp +++ b/mooncake-store/src/client_buffer.cpp @@ -149,6 +149,8 @@ uint64_t calculate_total_size(const Replica::Descriptor& replica) { total_length = disk_descriptor.object_size; } else if (replica.is_local_disk_replica()) { total_length = replica.get_local_disk_descriptor().object_size; + } else if (replica.is_dfs_replica()) { + total_length = replica.get_dfs_descriptor().object_size; } else if (replica.is_nof_replica()) { total_length = replica.get_nof_descriptor().buffer_descriptor.size_; } else { @@ -172,6 +174,9 @@ int allocateSlices(std::vector& slices, } else if (replica.is_local_disk_replica()) { slices.emplace_back( Slice{buffer_ptr, replica.get_local_disk_descriptor().object_size}); + } else if (replica.is_dfs_replica()) { + slices.emplace_back( + Slice{buffer_ptr, replica.get_dfs_descriptor().object_size}); } else if (replica.is_nof_replica()) { auto& handle = replica.get_nof_descriptor().buffer_descriptor; void* chunk_ptr = buffer_ptr; diff --git a/mooncake-store/src/client_service.cpp b/mooncake-store/src/client_service.cpp index 9e513d95a4..a6ff37f0ea 100644 --- a/mooncake-store/src/client_service.cpp +++ b/mooncake-store/src/client_service.cpp @@ -47,6 +47,7 @@ #endif #include "crc_checksum.h" #include "environ.h" +#include "storage/distributed/distributed_storage_backend.h" namespace mooncake { @@ -217,10 +218,13 @@ class ScatterRangeBuilder { struct ReplicaTransferSummary { size_t allocated_memory_replicas = 0; size_t allocated_nof_replicas = 0; + size_t allocated_dfs_replicas = 0; size_t successful_memory_transfers = 0; size_t successful_nof_transfers = 0; + size_t successful_dfs_transfers = 0; size_t failed_memory_transfers = 0; size_t failed_nof_transfers = 0; + size_t failed_dfs_transfers = 0; ErrorCode first_error = ErrorCode::OK; void RecordAllocatedReplica(const Replica::Descriptor& replica) { @@ -228,6 +232,8 @@ struct ReplicaTransferSummary { ++allocated_memory_replicas; } else if (replica.is_nof_replica()) { ++allocated_nof_replicas; + } else if (replica.is_dfs_replica()) { + ++allocated_dfs_replicas; } } @@ -236,6 +242,8 @@ struct ReplicaTransferSummary { ++successful_memory_transfers; } else if (replica_type == ReplicaType::NOF_SSD) { ++successful_nof_transfers; + } else if (replica_type == ReplicaType::DFS) { + ++successful_dfs_transfers; } } @@ -244,6 +252,8 @@ struct ReplicaTransferSummary { ++failed_memory_transfers; } else if (replica_type == ReplicaType::NOF_SSD) { ++failed_nof_transfers; + } else if (replica_type == ReplicaType::DFS) { + ++failed_dfs_transfers; } if (first_error == ErrorCode::OK) { first_error = error; @@ -251,9 +261,23 @@ struct ReplicaTransferSummary { } }; +bool NonDfsTransfersSucceeded(const ReplicaTransferSummary& summary) { + return summary.successful_memory_transfers == + summary.allocated_memory_replicas && + summary.successful_nof_transfers == summary.allocated_nof_replicas && + summary.failed_memory_transfers == 0 && + summary.failed_nof_transfers == 0; +} + +bool AllAllocatedTransfersSucceeded(const ReplicaTransferSummary& summary) { + return NonDfsTransfersSucceeded(summary) && + summary.successful_dfs_transfers == summary.allocated_dfs_replicas && + summary.failed_dfs_transfers == 0; +} + bool HasExpectedReplicaAllocation(const ReplicateConfig& config, const ReplicaTransferSummary& summary) { - if (config.nof_replica_num == 0) { + if (config.nof_replica_num == 0 && config.dfs_replica_num == 0) { return summary.allocated_memory_replicas > 0; } if (DetermineReplicaWriteMode(config) == @@ -263,7 +287,8 @@ bool HasExpectedReplicaAllocation(const ReplicateConfig& config, 0; } return summary.allocated_memory_replicas == config.replica_num && - summary.allocated_nof_replicas == config.nof_replica_num; + summary.allocated_nof_replicas == config.nof_replica_num && + summary.allocated_dfs_replicas == config.dfs_replica_num; } // success describes whether the overall put should succeed. Reliable modes @@ -285,12 +310,7 @@ FinalizeDecision DetermineFinalizeDecision( if (write_mode != ReplicaWriteMode::FLEXIBLE_DUAL_REPLICA) { const bool all_transfers_succeeded = - summary.successful_memory_transfers == - summary.allocated_memory_replicas && - summary.successful_nof_transfers == - summary.allocated_nof_replicas && - summary.failed_memory_transfers == 0 && - summary.failed_nof_transfers == 0; + AllAllocatedTransfersSucceeded(summary); if (allocation_satisfied && all_transfers_succeeded) { return {.end_type = ReplicaType::ALL, .revoke_type = std::nullopt, @@ -1323,7 +1343,11 @@ tl::expected Client::Get(const std::string& object_key, } auto t0_get = std::chrono::steady_clock::now(); - err = TransferRead(replica, slices); + if (replica.is_dfs_replica()) { + err = ReadDfsReplica(object_key, replica, slices); + } else { + err = TransferRead(replica, slices); + } // Release the cache block after transfer completes (memcpy is done) if (hot_cache_ && cache_used) { @@ -1589,6 +1613,8 @@ std::vector> Client::BatchGet( std::vector> pending_transfers; + std::vector dfs_read_requests; + std::vector dfs_read_indices; std::vector> results(object_keys.size()); // Record batch get transfer latency (Submit + Wait) auto t0_batch_get = std::chrono::steady_clock::now(); @@ -1630,7 +1656,18 @@ std::vector> Client::BatchGet( // Submit transfer operation asynchronously std::optional future; - if (replica.is_nof_replica()) { + if (replica.is_dfs_replica()) { + if (!dfs_storage_backend_) { + LOG(ERROR) << "DFS backend is not initialized"; + results[i] = tl::unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); + continue; + } + const auto& desc = replica.get_dfs_descriptor(); + dfs_read_requests.push_back( + DfsReadRequest{key, desc, slices_it->second}); + dfs_read_indices.push_back(i); + continue; + } else if (replica.is_nof_replica()) { auto contiguous_range = GetContiguousSliceRange(slices_it->second); if (!contiguous_range.has_value()) { LOG(ERROR) << "NoF transfer requires contiguous slices"; @@ -1662,6 +1699,35 @@ std::vector> Client::BatchGet( cache_used); } + if (!dfs_read_requests.empty()) { + auto dfs_results = dfs_storage_backend_->BatchRead(dfs_read_requests); + if (dfs_results.size() != dfs_read_requests.size()) { + LOG(ERROR) << "DFS BatchRead response size mismatch: expected " + << dfs_read_requests.size() << ", got " + << dfs_results.size(); + for (size_t index : dfs_read_indices) { + results[index] = tl::unexpected(ErrorCode::INTERNAL_ERROR); + } + } else { + for (size_t i = 0; i < dfs_results.size(); ++i) { + const size_t index = dfs_read_indices[i]; + const auto& request = dfs_read_requests[i]; + if (!dfs_results[i]) { + results[index] = tl::unexpected(dfs_results[i].error()); + continue; + } + auto checksum_result = VerifyObjectChecksum( + request.key, request.slices, request.descriptor.object_size, + query_results[index].object_checksum); + if (!checksum_result) { + results[index] = tl::unexpected(checksum_result.error()); + continue; + } + results[index] = {}; + } + } + } + // Wait for all transfers to complete for (auto& [index, key, future, stored_replica, cache_used] : pending_transfers) { @@ -1842,6 +1908,29 @@ tl::expected Client::Put(const ObjectKey& key, } } + if (transfer_summary.allocated_dfs_replicas > 0 && + NonDfsTransfersSucceeded(transfer_summary)) { + std::vector dfs_keys; + std::vector*> dfs_slices; + std::vector dfs_descriptors; + for (const auto& replica : start_result.value()) { + if (!replica.is_dfs_replica()) { + continue; + } + dfs_keys.push_back(key); + dfs_slices.push_back(&slices); + dfs_descriptors.push_back(replica.get_dfs_descriptor()); + } + for (auto dfs_result : + WriteDfsReplicas(dfs_keys, dfs_slices, dfs_descriptors)) { + if (dfs_result == ErrorCode::OK) { + transfer_summary.RecordSuccess(ReplicaType::DFS); + } else { + transfer_summary.RecordFailure(ReplicaType::DFS, dfs_result); + } + } + } + auto us_put = std::chrono::duration_cast( std::chrono::steady_clock::now() - t0_put) .count(); @@ -1921,6 +2010,11 @@ tl::expected Client::Upsert(const ObjectKey& key, return tl::unexpected(err); } + ReplicaTransferSummary transfer_summary; + for (const auto& replica : start_result.value()) { + transfer_summary.RecordAllocatedReplica(replica); + } + // Record transfer latency auto t0 = std::chrono::steady_clock::now(); @@ -1937,18 +2031,40 @@ tl::expected Client::Upsert(const ObjectKey& key, } } - // Transfer to memory replicas + // Transfer to memory and NoF replicas first. for (const auto& replica : start_result.value()) { - if (replica.is_memory_replica()) { + if (replica.is_memory_replica() || replica.is_nof_replica()) { + const auto replica_type = replica.is_memory_replica() + ? ReplicaType::MEMORY + : ReplicaType::NOF_SSD; ErrorCode transfer_err = TransferWrite(replica, slices); if (transfer_err != ErrorCode::OK) { - auto revoke_result = - master_client_.UpsertRevoke(key, ReplicaType::MEMORY); - if (!revoke_result) { - LOG(ERROR) << "Failed to revoke upsert operation"; - return tl::unexpected(revoke_result.error()); - } - return tl::unexpected(transfer_err); + transfer_summary.RecordFailure(replica_type, transfer_err); + continue; + } + transfer_summary.RecordSuccess(replica_type); + } + } + + if (transfer_summary.allocated_dfs_replicas > 0 && + NonDfsTransfersSucceeded(transfer_summary)) { + std::vector dfs_keys; + std::vector*> dfs_slices; + std::vector dfs_descriptors; + for (const auto& replica : start_result.value()) { + if (!replica.is_dfs_replica()) { + continue; + } + dfs_keys.push_back(key); + dfs_slices.push_back(&slices); + dfs_descriptors.push_back(replica.get_dfs_descriptor()); + } + for (auto dfs_result : + WriteDfsReplicas(dfs_keys, dfs_slices, dfs_descriptors)) { + if (dfs_result == ErrorCode::OK) { + transfer_summary.RecordSuccess(ReplicaType::DFS); + } else { + transfer_summary.RecordFailure(ReplicaType::DFS, dfs_result); } } } @@ -1960,13 +2076,27 @@ tl::expected Client::Upsert(const ObjectKey& key, metrics_->transfer_metric.put_latency_us.observe(us); } - // End upsert operation - auto end_result = master_client_.UpsertEnd(ObjectMeta{key, object_checksum}, - ReplicaType::MEMORY); - if (!end_result) { - ErrorCode err = end_result.error(); - LOG(ERROR) << "Failed to end upsert operation: " << err; - return tl::unexpected(err); + const auto finalize_decision = + DetermineFinalizeDecision(config, transfer_summary); + if (finalize_decision.end_type.has_value()) { + auto end_result = master_client_.UpsertEnd( + ObjectMeta{key, object_checksum}, *finalize_decision.end_type); + if (!end_result) { + ErrorCode err = end_result.error(); + LOG(ERROR) << "Failed to end upsert operation: " << err; + return tl::unexpected(err); + } + } + if (finalize_decision.revoke_type.has_value()) { + auto revoke_result = + master_client_.UpsertRevoke(key, *finalize_decision.revoke_type); + if (!revoke_result) { + LOG(ERROR) << "Failed to revoke upsert operation"; + return tl::unexpected(revoke_result.error()); + } + } + if (!finalize_decision.success) { + return tl::unexpected(finalize_decision.error); } // Success-side invalidation: a concurrent read between the pre-upsert @@ -2001,6 +2131,7 @@ std::vector> Client::BatchUpsert( auto t0 = std::chrono::steady_clock::now(); SubmitTransfers(ops); WaitForTransfers(ops); + SubmitDfsWrites(ops); auto us = std::chrono::duration_cast( std::chrono::steady_clock::now() - t0) .count(); @@ -2051,6 +2182,7 @@ class PutOperation { size_t requested_memory_replicas = 0; size_t requested_nof_replicas = 0; + size_t requested_dfs_replicas = 0; ReplicaTransferSummary transfer_summary; // Error context for debugging @@ -2103,12 +2235,14 @@ class PutOperation { void InitializeRequestedReplicas(const ReplicateConfig& config) { requested_memory_replicas = config.replica_num; requested_nof_replicas = config.nof_replica_num; + requested_dfs_replicas = config.dfs_replica_num; } ReplicateConfig ToReplicateConfig() const { ReplicateConfig config; config.replica_num = requested_memory_replicas; config.nof_replica_num = requested_nof_replicas; + config.dfs_replica_num = requested_dfs_replicas; return config; } @@ -2289,11 +2423,21 @@ void Client::StartBatchUpsert(std::vector& ops, // Process individual responses with robust error handling for (size_t i = 0; i < active_indices.size(); ++i) { auto& op = ops[active_indices[i]]; + op.InitializeRequestedReplicas(config); if (!start_responses[i]) { - op.SetError(start_responses[i].error(), - "Master failed to start upsert operation"); + op.SetTerminalError(start_responses[i].error(), + PutOperationState::MASTER_FAILED, + "Master failed to start upsert operation"); } else { op.replicas = start_responses[i].value(); + op.RecordAllocatedReplicas(); + if (!HasExpectedReplicaAllocation(config, op.transfer_summary)) { + op.SetTerminalError(ErrorCode::NO_AVAILABLE_HANDLE, + PutOperationState::MASTER_FAILED, + "Allocated replicas do not satisfy " + "requested replica policy"); + continue; + } VLOG(1) << "Successfully started upsert for key " << op.key << " with " << op.replicas.size() << " replicas"; } @@ -2424,6 +2568,137 @@ void Client::WaitForTransfers(std::vector& ops) { } } +std::vector Client::WriteDfsReplicas( + const std::vector& keys, + const std::vector*>& slice_lists, + const std::vector& descriptors) { + if (keys.size() != slice_lists.size() || + keys.size() != descriptors.size()) { + return std::vector(keys.size(), ErrorCode::INVALID_PARAMS); + } + if (keys.empty()) { + return {}; + } + if (!dfs_storage_backend_) { + LOG(ERROR) << "DFS backend is unavailable for synchronous write"; + return std::vector(keys.size(), + ErrorCode::DFS_SERVICE_UNAVAILABLE); + } + + std::vector results(keys.size(), ErrorCode::OK); + std::vector requests; + std::vector request_indices; + std::vector staging_buffers; + requests.reserve(keys.size()); + request_indices.reserve(keys.size()); + + auto runtime_accelerator = + device::GetAcceleratorRegistry().RuntimeAccelerators(); + for (size_t i = 0; i < keys.size(); ++i) { + if (slice_lists[i] == nullptr) { + results[i] = ErrorCode::INVALID_PARAMS; + continue; + } + + std::vector host_slices; + host_slices.reserve(slice_lists[i]->size()); + bool staging_succeeded = true; + for (const auto& slice : *slice_lists[i]) { + device::PointerInfo info{}; + auto* device = slice.ptr == nullptr + ? nullptr + : runtime_accelerator.FindDeviceForPointer( + slice.ptr, &info); + if (device == nullptr) { + host_slices.push_back(slice); + continue; + } + + device->SetContext(info.device_id); + auto buffer = pinned_buffer_pool_->Acquire(slice.size); + if (!device->Copy(buffer.data, slice.ptr, slice.size, + device::CopyDirection::kDeviceToHost)) { + LOG(ERROR) << "DFS D2H staging failed for key " << keys[i]; + pinned_buffer_pool_->Release(std::move(buffer)); + results[i] = ErrorCode::TRANSFER_FAIL; + staging_succeeded = false; + break; + } + host_slices.emplace_back(Slice{buffer.data, slice.size}); + staging_buffers.push_back(std::move(buffer)); + } + if (!staging_succeeded) { + continue; + } + + requests.push_back( + DfsWriteRequest{keys[i], descriptors[i], std::move(host_slices)}); + request_indices.push_back(i); + } + + auto write_results = dfs_storage_backend_->BatchWrite(requests); + if (write_results.size() != requests.size()) { + LOG(ERROR) << "DFS BatchWrite response size mismatch: expected " + << requests.size() << ", got " << write_results.size(); + for (size_t index : request_indices) { + results[index] = ErrorCode::INTERNAL_ERROR; + } + } else { + for (size_t i = 0; i < write_results.size(); ++i) { + results[request_indices[i]] = + write_results[i] ? ErrorCode::OK : write_results[i].error(); + } + } + + for (auto& buffer : staging_buffers) { + pinned_buffer_pool_->Release(std::move(buffer)); + } + return results; +} + +void Client::SubmitDfsWrites(std::vector& ops) { + std::vector keys; + std::vector*> slice_lists; + std::vector descriptors; + std::vector op_indices; + + for (size_t i = 0; i < ops.size(); ++i) { + auto& op = ops[i]; + if (op.IsResolved() || + op.transfer_summary.allocated_dfs_replicas == 0 || + !NonDfsTransfersSucceeded(op.transfer_summary)) { + continue; + } + + auto dfs_it = std::find_if(op.replicas.begin(), op.replicas.end(), + [](const Replica::Descriptor& replica) { + return replica.is_dfs_replica(); + }); + if (dfs_it == op.replicas.end()) { + op.transfer_summary.RecordFailure(ReplicaType::DFS, + ErrorCode::INVALID_REPLICA); + op.AppendFailureContext("Allocated DFS replica has no descriptor"); + continue; + } + keys.push_back(op.key); + slice_lists.push_back(&op.slices); + descriptors.push_back(dfs_it->get_dfs_descriptor()); + op_indices.push_back(i); + } + + auto results = WriteDfsReplicas(keys, slice_lists, descriptors); + for (size_t i = 0; i < results.size(); ++i) { + auto& op = ops[op_indices[i]]; + if (results[i] == ErrorCode::OK) { + op.transfer_summary.RecordSuccess(ReplicaType::DFS); + } else { + op.transfer_summary.RecordFailure(ReplicaType::DFS, results[i]); + op.AppendFailureContext("Synchronous DFS write failed: " + + toString(results[i])); + } + } +} + void Client::FinalizeBatchPut(std::vector& ops) { struct BatchFinalizeGroup { std::vector keys; @@ -2640,14 +2915,26 @@ void Client::FinalizeBatchUpsert(std::vector& ops) { for (size_t i = 0; i < ops.size(); ++i) { auto& op = ops[i]; - if (!op.IsResolved() && !op.replicas.empty() && - !op.pending_transfers.empty()) { + HasExpectedReplicaAllocation(op.ToReplicateConfig(), + op.transfer_summary) && + AllAllocatedTransfersSucceeded(op.transfer_summary)) { successful_object_metas.emplace_back( ObjectMeta{op.key, op.object_checksum}); successful_indices.emplace_back(i); - } else if (op.state != PutOperationState::PENDING && - !op.replicas.empty()) { + continue; + } + + if (!op.IsResolved() && !op.replicas.empty()) { + const auto error = op.transfer_summary.first_error == ErrorCode::OK + ? ErrorCode::TRANSFER_FAIL + : op.transfer_summary.first_error; + op.SetTerminalError( + error, PutOperationState::TRANSFER_FAILED, + op.failure_context.value_or( + "Replica transfer failed before upsert finalize")); + } + if (!op.IsSuccessful() && !op.replicas.empty()) { failed_keys.emplace_back(op.key); failed_indices.emplace_back(i); } @@ -2867,6 +3154,7 @@ std::vector> Client::BatchPutWhenPreferSameNode( seg_to_ops.at(seg).transfer_summary.first_error; op.failure_context = seg_to_ops.at(seg).failure_context; } + SubmitDfsWrites(ops); auto us = std::chrono::duration_cast( std::chrono::steady_clock::now() - t0) .count(); @@ -2908,6 +3196,7 @@ std::vector> Client::BatchPut( auto t0 = std::chrono::steady_clock::now(); SubmitTransfers(ops); WaitForTransfers(ops); + SubmitDfsWrites(ops); auto us = std::chrono::duration_cast( std::chrono::steady_clock::now() - t0) .count(); @@ -2975,7 +3264,6 @@ tl::expected Client::Remove(const ObjectKey& key, bool force) { if (!result) { return tl::unexpected(result.error()); } - if (hot_cache_) { hot_cache_->RemoveHotKey(key); } @@ -3463,6 +3751,12 @@ tl::expected Client::NotifyOffloadSuccess( return master_client_.NotifyOffloadSuccess(client_id_, tasks, metadatas); } +void Client::SetDfsStorageBackend( + std::shared_ptr backend) { + dfs_storage_backend_ = std::move(backend); + EnsureStorageControlPlaneStarted(); +} + tl::expected Client::PromotionObjectHeartbeat( std::vector& promotion_objects) { auto response = master_client_.PromotionObjectHeartbeat(client_id_); @@ -4105,6 +4399,8 @@ ErrorCode Client::TransferRead(const Replica::Descriptor& replica_descriptor, } else if (replica_descriptor.is_local_disk_replica()) { auto& disk_desc = replica_descriptor.get_local_disk_descriptor(); total_size = disk_desc.object_size; + } else if (replica_descriptor.is_dfs_replica()) { + total_size = replica_descriptor.get_dfs_descriptor().object_size; } size_t slices_size = CalculateSliceSize(slices); @@ -4114,9 +4410,40 @@ ErrorCode Client::TransferRead(const Replica::Descriptor& replica_descriptor, return ErrorCode::INVALID_PARAMS; } + if (replica_descriptor.is_dfs_replica()) { + LOG(ERROR) << "DFS reads require object key context"; + return ErrorCode::INVALID_REPLICA; + } + return TransferData(replica_descriptor, slices, TransferRequest::READ); } +ErrorCode Client::ReadDfsReplica(const std::string& key, + const Replica::Descriptor& replica_descriptor, + std::vector& slices) { + if (!replica_descriptor.is_dfs_replica()) { + return ErrorCode::INVALID_REPLICA; + } + if (!dfs_storage_backend_) { + LOG(ERROR) << "DFS backend is not initialized"; + return ErrorCode::DFS_SERVICE_UNAVAILABLE; + } + + const auto& desc = replica_descriptor.get_dfs_descriptor(); + std::vector requests{DfsReadRequest{key, desc, slices}}; + auto results = dfs_storage_backend_->BatchRead(requests); + if (results.size() != 1) { + LOG(ERROR) << "DFS BatchRead response size mismatch for key " << key; + return ErrorCode::INTERNAL_ERROR; + } + if (!results[0]) { + LOG(ERROR) << "DFS read failed for key " << key << ": " + << results[0].error(); + return results[0].error(); + } + return ErrorCode::OK; +} + ErrorCode Client::TransferReadRange( const Replica::Descriptor& replica_descriptor, std::vector& slices, uint64_t src_offset) { diff --git a/mooncake-store/src/file_storage.cpp b/mooncake-store/src/file_storage.cpp index 09ae610d5c..100c2e859f 100644 --- a/mooncake-store/src/file_storage.cpp +++ b/mooncake-store/src/file_storage.cpp @@ -12,6 +12,7 @@ #include "bool_parser.h" #include "environ.h" #include "storage_backend.h" +#include "storage/distributed/distributed_storage_backend.h" #include "client_metric.h" #include "utils.h" #include "device/accelerator_registry.h" @@ -98,6 +99,7 @@ FileStorageConfig FileStorageConfig::FromEnvironment() { config.storage_backend_type = StorageBackendType::kOffsetAllocator; } else if (storage_backend_descriptor == "distributed_storage_backend") { config.storage_backend_type = StorageBackendType::kDistributed; + config.enable_dfs = true; } else { LOG(ERROR) << "Unknown storage backend."; } @@ -272,6 +274,9 @@ FileStorage::FileStorage(const FileStorageConfig& config, pinned_buffer_pool_(std::make_unique()), client_buffer_allocator_(AlignedClientBufferAllocator::create( config.local_buffer_size, client ? client->GetProtocol() : "")) { + if (config_.storage_backend_type == StorageBackendType::kDistributed) { + config_.enable_dfs = true; + } if (!config.Validate()) { throw std::invalid_argument("Invalid FileStorage configuration"); } @@ -306,6 +311,13 @@ FileStorage::FileStorage(const FileStorageConfig& config, } storage_backend_ = create_storage_backend_result.value(); + if (auto distributed_backend = + std::dynamic_pointer_cast( + storage_backend_)) { + if (client_) { + client_->SetDfsStorageBackend(distributed_backend); + } + } // Register the client buffer with the process-wide io_uring fixed-buffer // mechanism. This must happen before any I/O threads start so that they @@ -354,6 +366,12 @@ tl::expected FileStorage::Init() { << init_storage_backend_result.error(); return init_storage_backend_result; } + if (config_.enable_dfs) { + client_buffer_gc_running_.store(true); + client_buffer_gc_thread_ = + std::thread(&FileStorage::ClientBufferGCThreadFunc, this); + return {}; + } auto enable_offloading_result = IsEnableOffloading(); if (enable_offloading_result.has_value()) { LOG(INFO) << "IsEnableOffloading result: " @@ -408,9 +426,15 @@ tl::expected FileStorage::Init() { }); if (!scan_meta_result) { - LOG(ERROR) << "Failed to scan meta and send to master: " - << scan_meta_result.error(); - return scan_meta_result; + if (config_.enable_dfs && + scan_meta_result.error() == ErrorCode::NOT_SUPPORTED) { + LOG(INFO) << "Currently, DFS backend does not support ScanMeta; " + "skip re-registering offloaded objects"; + } else { + LOG(ERROR) << "Failed to scan meta and send to master: " + << scan_meta_result.error(); + return scan_meta_result; + } } heartbeat_running_.store(true); @@ -879,8 +903,11 @@ tl::expected FileStorage::Heartbeat() { // === STEP 1: Send heartbeat and get offloading decisions === { MutexLocker locker(&offloading_mutex_); - auto heartbeat_result = client_->OffloadObjectHeartbeat( - enable_offloading_, offloading_objects); + auto fetch_offload_tasks = [&]() -> tl::expected { + return client_->OffloadObjectHeartbeat(enable_offloading_, + offloading_objects); + }; + auto heartbeat_result = fetch_offload_tasks(); if (!heartbeat_result) { ErrorCode err = heartbeat_result.error(); if (err == ErrorCode::SEGMENT_NOT_FOUND) { @@ -907,8 +934,7 @@ tl::expected FileStorage::Heartbeat() { << "heartbeat recovery: " << cap_result.error(); } } - heartbeat_result = client_->OffloadObjectHeartbeat( - enable_offloading_, offloading_objects); + heartbeat_result = fetch_offload_tasks(); if (!heartbeat_result) { LOG(ERROR) << "Heartbeat failed after re-registration: " << heartbeat_result.error(); @@ -1377,6 +1403,12 @@ tl::expected FileStorage::ReRegisterOffloadedObjects() { LOG(INFO) << "ReRegisterOffloadedObjects: ScanMeta returned. success=" << scan_meta_result.has_value(); if (!scan_meta_result) { + if (config_.enable_dfs && + scan_meta_result.error() == ErrorCode::NOT_SUPPORTED) { + LOG(INFO) << "ReRegisterOffloadedObjects: Currently, DFS ScanMeta " + "is not supported; skip"; + return {}; + } LOG(ERROR) << "ReRegisterOffloadedObjects: ScanMeta failed: " << scan_meta_result.error(); return scan_meta_result; diff --git a/mooncake-store/src/hf3fs/README.md b/mooncake-store/src/hf3fs/README.md index a542b05a9c..cf90411749 100644 --- a/mooncake-store/src/hf3fs/README.md +++ b/mooncake-store/src/hf3fs/README.md @@ -1,43 +1,70 @@ -# Mooncake HF3FS Plugin +# Mooncake HF3FS USRBIO Adapter -This plugin implements 3FS native API (USRBIO) as a high-performance storage backend for Mooncake. +> **Work in progress / experimental.** Descriptor-based DFS and this HF3FS +> adapter are intended for development and evaluation only. They are not +> production-ready or covered by Mooncake Store's general fault-tolerance, HA +> continuity, durability, or multi-tenant guarantees. + +This adapter implements the HF3FS native USRBIO data plane for Mooncake Store's +descriptor-based DFS replicas. It is selected explicitly with +`MOONCAKE_DFS_FS_ADAPTER=hf3fs`; the legacy `--root_fs_dir` option does not +enable it, and it does not automatically fall back to POSIX I/O. ## Prerequisites -### 1. 3FS Installation -- Build and install [3FS](https://github.com/deepseek-ai/3FS/) -- Required library: `libhf3fs_api_shared.so` (Default location: `3FS_PATH/build/src/lib/api`) - → Install to: `/usr/lib/` -- Required header: `hf3fs_usrbio.h` (Default location: `3FS_PATH/src/lib/api`) - → Install to: `/usr/include/` +### 1. HF3FS installation -### 2. Mooncake Configuration -- Enable 3FS support during CMake configuration: -```bash +- Build and install [3FS](https://github.com/deepseek-ai/3FS/). +- Install `libhf3fs_api_shared.so` from `3FS_PATH/build/src/lib/api` in a + library search path such as `/usr/lib/`. +- Install `hf3fs_usrbio.h` from `3FS_PATH/src/lib/api` in an include search + path such as `/usr/include/`. + +### 2. Mooncake build +Enable HF3FS support during CMake configuration: + +```bash cmake -DUSE_3FS=ON ... ``` -- Build and install Mooncake as usual. +Then build and install Mooncake as usual. ## Usage -### Basic Operation -Start master server and specify the 3FS mount point: +Configure the master with a shared HF3FS root and shard layout: + ```bash +export MOONCAKE_ENABLE_DFS=1 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true -./build/mooncake-store/src/mooncake_master \ - --root_fs_dir=/path/to/3fs_mount_point +./build/mooncake-store/src/mooncake_master [other master arguments] ``` -### Important Notes -1. The specified directory **must** be a 3FS mount point - - If not, the system will automatically fall back to POSIX API -2. For optimal performance: - - Ensure proper permissions on the 3FS mount point - - Verify 3FS service is running before execution - -### Example + +Every client that may read or write DFS replicas must initialize FileStorage's +distributed backend with the same absolute root path and layout. Use the same +root path string in every process: + ```bash +export MOONCAKE_OFFLOAD_ENABLED=true +export MOONCAKE_OFFLOAD_STORAGE_BACKEND_DESCRIPTOR=distributed_storage_backend +export MOONCAKE_OFFLOAD_FILE_STORAGE_PATH=/data/file_storage +export MOONCAKE_MASTER=127.0.0.1:50051 +export MOONCAKE_DFS_ROOT_DIR=/mnt/3fs/mooncake +export MOONCAKE_DFS_FS_ADAPTER=hf3fs +export MOONCAKE_DFS_SHARD_COUNT=64 +export MOONCAKE_DFS_SHARD_CAPACITY=4294967296 +export MOONCAKE_DFS_ALIGNMENT=4096 +export MOONCAKE_DFS_SINGLE_TENANT=true + +python -m mooncake.mooncake_store_service +``` -ROLE=prefill MOONCAKE_STORAGE_ROOT_DIR=/mnt/3fs python3 ./stress_cluster_benchmark.py -``` \ No newline at end of file +`MOONCAKE_OFFLOAD_FILE_STORAGE_PATH` must already be an absolute, writable, +non-symlink directory. For the full configuration and current limitations, see +the [DFS deployment documentation](../../../docs/source/deployment/mooncake-store-deployment-guide.md#dfs-storage). diff --git a/mooncake-store/src/master_service.cpp b/mooncake-store/src/master_service.cpp index 2bd9d0bc0f..c03412bfbf 100644 --- a/mooncake-store/src/master_service.cpp +++ b/mooncake-store/src/master_service.cpp @@ -11,10 +11,12 @@ #include #include #include +#include #include #include #include #include +#include #include #include #include @@ -26,6 +28,7 @@ #include "http_metadata_server.h" #include "master_metric_manager.h" #include "common.h" +#include "environ.h" #include "segment.h" #ifdef USE_HTTP #include "transfer_metadata_plugin.h" @@ -48,6 +51,8 @@ #include "ha/snapshot/snapshot_logger.h" #include "utils/zstd_util.h" #include "utils/file_util.h" +#include "storage/distributed/dfs_global_allocator.h" +#include "storage/distributed/distributed_storage_backend.h" #include "random.h" #include "utils.h" #include "kv_event/kv_event_config.h" @@ -122,7 +127,7 @@ uint64_t SaturatingMultiply(uint64_t lhs, uint64_t rhs) { bool HasExpectedReplicaAllocation(const ReplicateConfig& config, size_t allocated_memory_replicas, size_t allocated_nof_replicas) { - if (config.nof_replica_num == 0) { + if (config.nof_replica_num == 0 && config.dfs_replica_num == 0) { return allocated_memory_replicas > 0; } if (DetermineReplicaWriteMode(config) == @@ -399,6 +404,7 @@ MasterService::MasterService(const MasterServiceConfig& config) << ")"; } + InitDfsAllocatorFromEnvironment(config); kv_event_publisher_ = std::make_unique(BuildKvEventConfig(config)); @@ -511,6 +517,43 @@ MasterService::MasterService(const MasterServiceConfig& config) } } +void MasterService::InitDfsAllocatorFromEnvironment( + const MasterServiceConfig& config) { + enable_dfs_ = Environ::GetBool( + "MOONCAKE_ENABLE_DFS", Environ::GetBool("MOONCAKE_DFS_ENABLED", false)); + if (!enable_dfs_) return; + + if (config.enable_snapshot || config.enable_snapshot_restore || + enable_oplog_) { + LOG(ERROR) << "DFS cannot be enabled with snapshot or oplog recovery " + "until DFS allocator state restoration is supported"; + throw std::invalid_argument( + "DFS is incompatible with snapshot/oplog recovery"); + } + + const auto dfs_config = DistributedStorageConfig::FromEnvironment(); + if (!dfs_config.single_tenant) { + LOG(ERROR) << "Currently, DFS backend is not supported in " + "multi-tenant mode"; + enable_dfs_ = false; + return; + } + + dfs_allocator_ = std::make_unique(); + auto init_result = dfs_allocator_->Init(dfs_config); + if (!init_result) { + LOG(ERROR) << "Failed to initialize DFS allocator, error=" + << init_result.error() << ", config={" + << dfs_config.FormatStr() << "}"; + dfs_allocator_.reset(); + enable_dfs_ = false; + return; + } + + LOG(INFO) << "DFS allocator initialized, config={" << dfs_config.FormatStr() + << "}"; +} + std::unique_ptr MasterService::CreateSnapshotCatalogStore(const MasterServiceConfig& config) { auto catalog_kind = @@ -636,6 +679,8 @@ void MasterService::RunNoFBatchEvictForTesting(double evict_ratio_target, NoFBatchEvict(evict_ratio_target, evict_ratio_lowerbound); } +void MasterService::RunDfsEvictionForTesting() { RunDfsEviction(); } + void MasterService::SetNoFProbeFnForTesting(NoFProbeFn fn) { #ifdef USE_NOF std::lock_guard lock(nof_probe_fn_mutex_); @@ -1647,6 +1692,7 @@ size_t MasterService::EraseReplicasWithCacheTotalAccounting( // Release SSD/local-disk usage for any local-disk replicas being removed. // No-op for memory/noF replicas, so it is safe to call unconditionally. ReleaseLocalDiskUsage(erased_replicas); + FreeDfsReplicas(metadata.user_key, erased_replicas); return erased_replicas.size(); } @@ -1748,6 +1794,7 @@ void MasterService::FinalizeRemovedReplicasAfterDurable( erased_replicas.begin(), erased_replicas.end(), [](const Replica& replica) { return replica.is_local_disk_replica(); }); ReleaseLocalDiskUsage(erased_replicas); + FreeDfsReplicas(metadata.user_key, erased_replicas); if (erased_local_disk) { shard.OnDiskReplicaRemoved(erased_local_disk, metadata); } @@ -1804,6 +1851,7 @@ void MasterService::FinalizeExpiredProcessingReplicasAfterDurable( auto replicas = PopReplicasWithCacheTotalAccounting( metadata, &Replica::fn_is_processing); if (!replicas.empty()) { + FreeDfsReplicas(metadata.user_key, replicas); std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back(std::move(replicas), ttl); } @@ -1844,6 +1892,7 @@ void MasterService::FinalizeExpiredReplicationTaskAfterDurable( metadata, [&ids](const Replica& replica) { return ids.contains(replica.id()); }); if (!replicas.empty()) { + FreeDfsReplicas(metadata.user_key, replicas); std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back(std::move(replicas), ttl); } @@ -1993,6 +2042,7 @@ MasterService::EraseMetadata( ErasePromotionTaskIfPresent(tenant_state, key); ReleaseLocalDiskUsage(metadata.GetAllReplicas()); + FreeDfsReplicas(key, metadata.GetAllReplicas()); AccountCacheTotalRemoval(metadata); if (metadata.GetCommittedSoftPinTimeout()) { soft_pin_deadline_index_.Remove(tenant_id.MakeScopedKey(key)); @@ -2797,6 +2847,10 @@ void MasterService::RestoreFromStandbySnapshot( const std::vector& objects, uint64_t initial_oplog_sequence_id, const std::vector& segments) { + if (enable_dfs_) { + throw std::runtime_error( + "DFS standby restore requires DFS allocator state restoration"); + } // The ordered writer initializes its sequence from durable_prefix. (void)initial_oplog_sequence_id; @@ -3314,8 +3368,13 @@ auto MasterService::GetReplicaList(const std::string& key, [this](const Replica& replica) { return IsReplicaReadable(replica); }, - [&replica_list](const Replica& replica) { + [this, &key, &replica_list](const Replica& replica) { replica_list.emplace_back(replica.get_descriptor()); + if (replica.is_dfs_replica() && dfs_allocator_) { + const auto& desc = replica.get_dfs_descriptor(); + dfs_allocator_->UpdateAccess(key, desc.shard_idx, + desc.offset); + } }); if (replica_list.empty()) { @@ -3823,6 +3882,23 @@ auto MasterService::AllocateAndInsertMetadata( ReplicaStatus::PROCESSING); } + if (config.dfs_replica_num > 0) { + if (!dfs_allocator_ || !dfs_allocator_->IsInitialized()) { + LOG(ERROR) << "Failed to allocate DFS replica for key=" << key + << ", error=dfs_allocator_not_initialized"; + refund_pending_quota(); + return tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); + } + auto alloc = dfs_allocator_->Allocate(key, value_length); + if (!alloc) { + LOG(ERROR) << "Failed to allocate DFS replica for key=" << key + << ", error=" << alloc.error(); + refund_pending_quota(); + return tl::make_unexpected(alloc.error()); + } + replicas.emplace_back(std::move(*alloc), ReplicaStatus::PROCESSING); + } + std::vector replica_list; std::vector eligible_replica_ids; replica_list.reserve(replicas.size()); @@ -3847,6 +3923,11 @@ auto MasterService::AllocateAndInsertMetadata( << nof_desc.buffer_descriptor.buffer_address_ << ", transport_endpoint=" << nof_desc.buffer_descriptor.transport_endpoint_; + } else if (replica.is_dfs_replica()) { + const auto& dfs_desc = desc.get_dfs_descriptor(); + VLOG(1) << "Replica #" << ++i << ": dfs_file=" << dfs_desc.file_path + << ", offset=" << dfs_desc.offset + << ", shard_idx=" << dfs_desc.shard_idx; } } @@ -3857,6 +3938,7 @@ auto MasterService::AllocateAndInsertMetadata( config.with_hard_pin, config.data_type, group_id, tenant_id, key)); if (!inserted) { + FreeDfsReplicas(key, replicas); LOG(INFO) << "key=" << key << ", info=object_already_exists"; refund_pending_quota(); return tl::make_unexpected(ErrorCode::OBJECT_ALREADY_EXISTS); @@ -3891,14 +3973,33 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, const ReplicateConfig& config) -> tl::expected, ErrorCode> { const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); - if ((config.replica_num == 0 && config.nof_replica_num == 0) || + if ((config.replica_num == 0 && config.nof_replica_num == 0 && + config.dfs_replica_num == 0) || key.empty() || slice_length == 0) { LOG(ERROR) << "key=" << key << ", replica_num=" << config.replica_num << ", nof_replica_num=" << config.nof_replica_num + << ", dfs_replica_num=" << config.dfs_replica_num << ", slice_length=" << slice_length << ", key_size=" << key.size() << ", error=invalid_params"; return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } + if (config.dfs_replica_num > 1 || + (config.dfs_replica_num > 0 && config.replica_num == 0)) { + LOG(ERROR) << "key=" << key << ", replica_num=" << config.replica_num + << ", dfs_replica_num=" << config.dfs_replica_num + << ", error=invalid_dfs_replica_config"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + if (config.dfs_replica_num > 0 && !object_id.tenant_id.IsDefault()) { + LOG(ERROR) << "key=" << key << ", tenant_id=" << tenant_id + << ", error=dfs_currently_requires_default_tenant"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + if (config.dfs_replica_num > 0 && + (!enable_dfs_ || !dfs_allocator_ || !dfs_allocator_->IsInitialized())) { + LOG(ERROR) << "key=" << key << ", error=dfs_allocator_not_initialized"; + return tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); + } if (config.prefer_alloc_in_same_node && config.nof_replica_num > 0) { LOG(ERROR) << "key=" << key << ", nof_replica_num=" << config.nof_replica_num @@ -4009,6 +4110,7 @@ auto MasterService::PutStart(const UUID& client_id, const std::string& key, auto replicas = PopReplicasWithCacheTotalAccounting( metadata, &Replica::fn_is_processing); if (!replicas.empty()) { + FreeDfsReplicas(key, replicas); std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back( std::move(replicas), @@ -4096,7 +4198,8 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, return (replica.is_memory_replica() && !replica.has_invalid_mem_handle()) || (replica.is_nof_replica() && - !replica.has_invalid_nof_handle()); + !replica.has_invalid_nof_handle()) || + replica.is_dfs_replica(); } if (replica_type == ReplicaType::MEMORY) { return replica.is_memory_replica() && @@ -4140,12 +4243,16 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, [&is_target_replica](const Replica& replica) { return replica.is_processing() && is_target_replica(replica); }, - [&metadata, &completed_pending_replica](Replica& replica) { + [this, &key, &metadata, &completed_pending_replica](Replica& replica) { if (replica.is_processing() && metadata.PendingSoftPinOwnsReplica(replica.id())) { completed_pending_replica = true; } replica.mark_complete(); + if (replica.is_dfs_replica() && dfs_allocator_) { + const auto& desc = replica.get_dfs_descriptor(); + dfs_allocator_->UpdateAccess(key, desc.shard_idx, desc.offset); + } }); if (!had_completed_replica && completed_pending_replica && @@ -4168,7 +4275,10 @@ auto MasterService::PutEnd(const UUID& client_id, const ObjectMeta& object_meta, return tl::make_unexpected(settle_result.error()); } - if (enable_offload_ && !offload_on_evict_) { + if (replica_type != ReplicaType::DFS && enable_offload_ && + !offload_on_evict_ && !metadata.HasReplica([](const Replica& replica) { + return replica.is_dfs_replica() && replica.is_processing(); + })) { auto& tenant_state = accessor.GetTenantState(); bool task_created = false; metadata.VisitReplicas( @@ -4341,14 +4451,15 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, return tl::make_unexpected(ErrorCode::INVALID_WRITE); } - auto processing_rep = metadata.GetFirstReplica([replica_type]( - const Replica& replica) { - if (replica_type == ReplicaType::ALL) { - return (replica.is_memory_replica() || replica.is_nof_replica()) && - !replica.is_processing(); - } - return replica.type() == replica_type && !replica.is_processing(); - }); + auto processing_rep = + metadata.GetFirstReplica([replica_type](const Replica& replica) { + if (replica_type == ReplicaType::ALL) { + return (replica.is_memory_replica() || + replica.is_nof_replica() || replica.is_dfs_replica()) && + !replica.is_processing(); + } + return replica.type() == replica_type && !replica.is_processing(); + }); if (processing_rep != nullptr) { LOG(ERROR) << "key=" << key << ", status=" << processing_rep->status() << ", error=invalid_replica_status"; @@ -4360,7 +4471,8 @@ auto MasterService::PutRevoke(const UUID& client_id, const std::string& key, return false; } if (replica_type == ReplicaType::ALL) { - return r.is_memory_replica() || r.is_nof_replica(); + return r.is_memory_replica() || r.is_nof_replica() || + r.is_dfs_replica(); } return r.type() == replica_type; }; @@ -4496,14 +4608,36 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, -> tl::expected, ErrorCode> { const auto object_id = MakeObjectIdentityForRequest(key, tenant_id); // --- Parameter validation (same as PutStart) --- - if ((config.replica_num == 0 && config.nof_replica_num == 0) || + if ((config.replica_num == 0 && config.nof_replica_num == 0 && + config.dfs_replica_num == 0) || key.empty() || slice_length == 0) { LOG(ERROR) << "key=" << key << ", replica_num=" << config.replica_num << ", nof_replica_num=" << config.nof_replica_num + << ", dfs_replica_num=" << config.dfs_replica_num << ", slice_length=" << slice_length << ", key_size=" << key.size() << ", error=invalid_params"; return tl::make_unexpected(ErrorCode::INVALID_PARAMS); } + if (config.dfs_replica_num > 1 || + (config.dfs_replica_num > 0 && config.replica_num == 0)) { + LOG(ERROR) << "key=" << key << ", replica_num=" << config.replica_num + << ", dfs_replica_num=" << config.dfs_replica_num + << ", error=invalid_dfs_replica_config"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + if (config.dfs_replica_num > 0 && !object_id.tenant_id.IsDefault()) { + LOG(ERROR) << "key=" << key << ", tenant_id=" << tenant_id + << ", error=dfs_currently_requires_default_tenant"; + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + if (config.dfs_replica_num > 0 && + (!enable_dfs_ || dfs_allocator_ == nullptr || + !dfs_allocator_->IsInitialized())) { + LOG(ERROR) << "key=" << key + << ", dfs_replica_num=" << config.dfs_replica_num + << ", error=dfs_service_unavailable"; + return tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); + } if (config.prefer_alloc_in_same_node && config.nof_replica_num > 0) { LOG(ERROR) << "key=" << key << ", nof_replica_num=" << config.nof_replica_num @@ -4653,6 +4787,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, metadata.PopReplicas(&Replica::fn_is_processing); metadata.ClearPendingSoftPinAction(); if (!processing_replicas.empty()) { + FreeDfsReplicas(key, processing_replicas); std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back( std::move(processing_replicas), @@ -4725,6 +4860,33 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, // hard_pinned is const and preserved automatically — upsert // does not change the eviction protection level of an // existing object. + const size_t existing_dfs_replicas = + metadata.CountReplicas(&Replica::fn_is_dfs_replica); + if (config.dfs_replica_num > 0 || + existing_dfs_replicas > 0) { + const size_t existing_memory_replicas = + metadata.CountReplicas( + &Replica::fn_is_memory_replica); + const size_t existing_nof_replicas = + metadata.CountReplicas(&Replica::fn_is_nof_replica); + if (existing_memory_replicas != config.replica_num || + existing_nof_replicas != config.nof_replica_num || + existing_dfs_replicas != config.dfs_replica_num) { + LOG(ERROR) + << "key=" << key + << ", error=dfs_upsert_topology_mismatch" + << ", existing_memory=" + << existing_memory_replicas + << ", requested_memory=" << config.replica_num + << ", existing_nof=" << existing_nof_replicas + << ", requested_nof=" << config.nof_replica_num + << ", existing_dfs=" << existing_dfs_replicas + << ", requested_dfs=" << config.dfs_replica_num; + return tl::make_unexpected( + ErrorCode::INVALID_PARAMS); + } + } + metadata.client_id = client_id; metadata.put_start_time = now; @@ -4798,6 +4960,7 @@ auto MasterService::UpsertStart(const UUID& client_id, const std::string& key, auto old_replicas = PopReplicasWithCacheTotalAccounting(metadata); if (!old_replicas.empty()) { + FreeDfsReplicas(key, old_replicas); std::lock_guard lock(discarded_replicas_mutex_); discarded_replicas_.emplace_back( std::move(old_replicas), @@ -5699,6 +5862,7 @@ tl::expected MasterService::MoveEnd( return replica.id() == source_id; }); if (!source_replica.empty()) { + FreeDfsReplicas(key, source_replica); if (enable_multi_tenants_) { auto release_result = metadata.quota_ledger.ReleaseCommitted( GetBoundTenantQuotaHandle(accessor.GetTenantState()), @@ -6381,6 +6545,153 @@ bool MasterService::CleanupStaleHandles( return !metadata.IsValid(); } +void MasterService::FreeDfsReplicas(const std::string& key, + const std::vector& replicas) { + if (!dfs_allocator_) return; + for (const auto& replica : replicas) { + if (!replica.is_dfs_replica()) continue; + const auto& desc = replica.get_dfs_descriptor(); + dfs_allocator_->Free(desc.offset, desc.aligned_size, desc.shard_idx, + key); + } +} + +void MasterService::RunDfsEviction() { + if (!dfs_allocator_) return; + + const TenantId tenant_id = TenantId::Default(); + using CandidateIdentity = std::tuple; + std::set attempted; + + while (true) { + auto pending = dfs_allocator_->PrepareEviction(); + if (pending.Empty()) return; + + const auto candidates = pending.Candidates(); + std::vector accepted(candidates.size(), false); + std::vector considered(candidates.size(), false); + bool saw_repeated_candidate = false; + + // Group prepared candidates by metadata shard, then validate and + // remove each group while holding only that shard. Prepared allocator + // extents remain unavailable until ResolvePreparedEviction(), so + // different metadata shards do not need one cross-shard transaction. + std::array, kNumShards> indexes_by_shard; + for (size_t i = 0; i < candidates.size(); ++i) { + const auto& candidate = candidates[i]; + if (!attempted + .emplace(candidate.key, candidate.shard_idx, + candidate.offset) + .second) { + saw_repeated_candidate = true; + continue; + } + considered[i] = true; + indexes_by_shard[getMetadataShardIndex(tenant_id, candidate.key)] + .push_back(i); + } + + auto matches_candidate = [](const Replica& replica, + const auto& candidate) { + return replica.is_dfs_replica() && + replica.get_dfs_descriptor().shard_idx == + candidate.shard_idx && + replica.get_dfs_descriptor().offset == candidate.offset; + }; + + auto now = std::chrono::system_clock::now(); + for (size_t shard_idx = 0; shard_idx < kNumShards; ++shard_idx) { + if (indexes_by_shard[shard_idx].empty()) continue; + + std::shared_lock snapshot_lock(snapshot_mutex_); + SharedMutexLocker shard_lock(&metadata_shards_[shard_idx].mutex); + + // Validate and remove candidates under the same shard lock. Once a + // candidate has been seen in this cycle, it is excluded above; + // encountering it again means the LRU scan has wrapped. + for (const size_t i : indexes_by_shard[shard_idx]) { + const auto& candidate = candidates[i]; + auto tenant_it = + metadata_shards_[shard_idx].tenants.find(tenant_id); + if (tenant_it == metadata_shards_[shard_idx].tenants.end()) { + accepted[i] = true; + continue; + } + auto& tenant_state = tenant_it->second; + auto metadata_it = tenant_state.metadata.find(candidate.key); + if (metadata_it == tenant_state.metadata.end()) { + accepted[i] = true; + continue; + } + + auto& metadata = metadata_it->second; + const bool has_candidate = + metadata.HasReplica([&](const Replica& replica) { + return matches_candidate(replica, candidate); + }); + if (!has_candidate) { + accepted[i] = true; + continue; + } + + const bool candidate_is_processing = + metadata.HasReplica([&](const Replica& replica) { + return matches_candidate(replica, candidate) && + replica.is_processing(); + }); + accepted[i] = + !candidate_is_processing && + !tenant_state.processing_keys.contains(candidate.key) && + !metadata.IsHardPinned() && metadata.IsLeaseExpired(now) && + (!IsSoftPinActive(metadata, now) || + allow_evict_soft_pinned_objects_); + if (!accepted[i]) continue; + + // A missing descriptor was accepted above and is already + // evicted from the master's point of view. Otherwise remove + // the accepted descriptor before its allocator extent can be + // committed and reused. + const size_t erased = + metadata.EraseReplicas([&](const Replica& replica) { + return matches_candidate(replica, candidate) && + !replica.is_processing(); + }); + if (erased > 0 && !metadata.IsValid()) { + PublishKvRemovedAfterEvict(candidate.key, + metadata.size * erased, "disk", + metadata, tenant_id); + EraseMetadata(tenant_state, metadata_it, tenant_id, + QuotaEraseMode::kFull); + } + } + + auto tenant_it = + metadata_shards_[shard_idx].tenants.find(tenant_id); + if (tenant_it != metadata_shards_[shard_idx].tenants.end() && + tenant_it->second.Empty()) { + metadata_shards_[shard_idx].tenants.erase(tenant_it); + } + } + + dfs_allocator_->ResolvePreparedEviction(std::move(pending), accepted); + + // A protected allocation stays live, but moving it to the MRU side + // lets this cycle inspect colder candidates behind it. The attempted + // set bounds the scan when every remaining allocation is protected. + for (size_t i = 0; i < candidates.size(); ++i) { + if (considered[i] && !accepted[i]) { + dfs_allocator_->UpdateAccess(candidates[i].key, + candidates[i].shard_idx, + candidates[i].offset); + } + } + + if (saw_repeated_candidate) { + return; + } + } +} + size_t MasterService::GetKeyCount() const { size_t total = 0; for (size_t i = 0; i < kNumShards; i++) { @@ -7621,6 +7932,7 @@ void MasterService::EvictionThreadFunc() { VLOG(1) << "action=eviction_thread_started"; auto last_discard_time = std::chrono::system_clock::now(); + auto next_dfs_eviction_time = std::chrono::steady_clock::now(); while (eviction_running_) { const auto now = std::chrono::system_clock::now(); double used_ratio = @@ -7671,6 +7983,16 @@ void MasterService::EvictionThreadFunc() { } #endif + if (dfs_allocator_ && dfs_allocator_->IsEvictionEnabled()) { + const auto steady_now = std::chrono::steady_clock::now(); + if (steady_now >= next_dfs_eviction_time) { + RunDfsEviction(); + next_dfs_eviction_time = + std::chrono::steady_clock::now() + + dfs_allocator_->GetEvictionCheckInterval(); + } + } + if (promotion_candidate_count_.load(std::memory_order_relaxed) > 0) { RunPromotionCandidateRetry(); } @@ -7780,6 +8102,7 @@ void MasterService::DiscardExpiredProcessingReplicas( metadata.PopReplicas(&Replica::fn_is_processing); metadata.ClearPendingSoftPinIfNoViableReplica(); if (!replicas.empty()) { + FreeDfsReplicas(*key_it, replicas); discarded_replicas.emplace_back(std::move(replicas), ttl); } if (!metadata.IsValid()) { @@ -7895,6 +8218,7 @@ void MasterService::DiscardExpiredProcessingReplicas( auto replicas = PopReplicasWithCacheTotalAccounting(metadata, target_pred); if (!replicas.empty()) { + FreeDfsReplicas(task_it->first, replicas); discarded_replicas.emplace_back(std::move(replicas), ttl); } if (!metadata.IsValid()) { diff --git a/mooncake-store/src/real_client.cpp b/mooncake-store/src/real_client.cpp index df68476af8..a8fda981c6 100644 --- a/mooncake-store/src/real_client.cpp +++ b/mooncake-store/src/real_client.cpp @@ -2683,10 +2683,10 @@ std::shared_ptr RealClient::get_buffer_internal( return nullptr; } - // Select best replica: prefer local MEMORY, then any MEMORY, - // then LOCAL_DISK, then DISK. + // Select best replica: prefer local MEMORY, any MEMORY, local NOF, any NOF, + // LOCAL_DISK, DFS, then DISK. // LOCAL_DISK data is on a remote node's SSD — must use offload RPC. - // MEMORY / DISK are handled via client_->Get below. + // MEMORY / DISK / DFS are handled via client_->Get below. auto local_endpoints = client_->GetLocalEndpoints(); const auto *best_replica = SelectBestReplica(replica_list, local_endpoints); if (!best_replica) { @@ -2736,14 +2736,15 @@ std::shared_ptr RealClient::get_buffer_internal( return buffer_handle; } - // MEMORY / DISK: use client_->Get. FilterQueryResult ensures + // MEMORY / DISK / DFS: use client_->Get. FilterQueryResult ensures // Client::Get internal FindFirstCompleteReplica can only see // the replica we selected, preventing accidental LOCAL_DISK picks. auto runtime_accelerator = device::GetAcceleratorRegistry().RuntimeAccelerators(); - if (replica.is_disk_replica() && + if ((replica.is_disk_replica() || replica.is_dfs_replica()) && runtime_accelerator.FindDeviceForPointer(buffer_handle->ptr())) { - LOG(WARNING) << "DISK replica for key '" << key + LOG(WARNING) << (replica.is_dfs_replica() ? "DFS" : "DISK") + << " replica for key '" << key << "' received a device pointer from the allocator; " << "file I/O cannot write to GPU memory — read will fail. " << "Ensure client_buffer_allocator_ returns host memory."; @@ -3039,8 +3040,8 @@ RealClient::batch_get_buffer_internal( continue; } - // Select best replica: prefer local MEMORY, then any MEMORY, - // then LOCAL_DISK, then DISK. + // Select best replica: prefer local MEMORY, any MEMORY, local NOF, + // any NOF, LOCAL_DISK, DFS, then DISK. const auto *best_replica = SelectBestReplica(query_result_values.replicas, local_endpoints); if (!best_replica) { @@ -3078,15 +3079,16 @@ RealClient::batch_get_buffer_internal( continue; } - // DISK replicas use storage_backend::vector_read (file I/O) which - // can only write to CPU-addressable memory. If the allocator ever - // returns device memory for DISK, the read will silently fail. + // File-backed replicas use file I/O which can only write to + // CPU-addressable memory. If the allocator ever returns device memory, + // the read will fail. auto runtime_accelerator = device::GetAcceleratorRegistry().RuntimeAccelerators(); - if (replica.is_disk_replica() && + if ((replica.is_disk_replica() || replica.is_dfs_replica()) && runtime_accelerator.FindDeviceForPointer(buffer_handle->ptr())) { LOG(WARNING) - << "DISK replica for key '" << key + << (replica.is_dfs_replica() ? "DFS" : "DISK") + << " replica for key '" << key << "' received a device pointer from the allocator; " << "file I/O cannot write to GPU memory — read will fail. " << "Ensure client_buffer_allocator_ returns host memory."; @@ -3357,13 +3359,16 @@ tl::expected RealClient::execute_ranged_read( return static_cast(total_size); } - if (replica.is_disk_replica()) { - // DISK full read: local file I/O (vector_read) cannot write to - // GPU memory. Use temp CPU buffer, then scatter to dst. + if (replica.is_disk_replica() || replica.is_dfs_replica()) { + // File-backed full read: use a contiguous temp buffer, then + // scatter to the requested destination. + const char *replica_type = + replica.is_dfs_replica() ? "DFS" : "DISK"; auto alloc_result = client_buffer_allocator_->allocate(total_size); if (!alloc_result) { - LOG(ERROR) << "Failed to allocate temp buffer for DISK full " - << "read, key: " << key << ", size: " << total_size; + LOG(ERROR) << "Failed to allocate temp buffer for " + << replica_type << " full read, key: " << key + << ", size: " << total_size; return tl::unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } BufferHandle tmp_handle(std::move(*alloc_result)); @@ -3373,14 +3378,15 @@ tl::expected RealClient::execute_ranged_read( FilterQueryResult(query_result, replica, verify_checksum); auto get_result = client_->Get(key, filtered_qr, tmp_slices); if (!get_result) { - LOG(ERROR) << "DISK Get failed for key: " << key + LOG(ERROR) << replica_type << " Get failed for key: " << key << " with error: " << toString(get_result.error()); return tl::unexpected(get_result.error()); } void *dst = static_cast(buffer) + dst_offset; const void *src = tmp_handle.ptr(); if (auto r = scatter_host_to_maybe_device( - dst, src, total_size, "DISK full read, key: " + key); + dst, src, total_size, + std::string(replica_type) + " full read, key: " + key); !r) { return tl::unexpected(r.error()); } @@ -3438,10 +3444,10 @@ tl::expected RealClient::execute_ranged_read( return static_cast(total_size); } - // Partial disk read: allocate temp CPU buffer, invoke read_op to + // Partial file-backed read: allocate temp buffer, invoke read_op to // fill it, then scatter [src_offset, src_offset+size) to dst. // - // buf_size controls how much to allocate / read. DISK must use + // buf_size controls how much to allocate / read. DISK/DFS must use // total_size (allocateSlices requires full-object slices). // LOCAL_DISK can use src_offset + size (offload RPC transfers // sequentially from remote offset 0). @@ -3450,7 +3456,8 @@ tl::expected RealClient::execute_ranged_read( size_t buf_size) -> tl::expected { auto alloc_result = client_buffer_allocator_->allocate(buf_size); if (!alloc_result) { - LOG(ERROR) << "Failed to allocate temp buffer for ranged disk " + LOG(ERROR) << "Failed to allocate temp buffer for ranged " + << "file-backed " << "read, key: " << key << ", size: " << buf_size; return tl::unexpected(ErrorCode::NO_AVAILABLE_HANDLE); } @@ -3461,7 +3468,7 @@ tl::expected RealClient::execute_ranged_read( const void *src = static_cast(tmp_handle.ptr()) + src_offset; if (auto r = scatter_host_to_maybe_device( - dst, src, size, "ranged disk read, key: " + key); + dst, src, size, "ranged file-backed read, key: " + key); !r) { return tl::unexpected(r.error()); } @@ -3500,9 +3507,10 @@ tl::expected RealClient::execute_ranged_read( src_offset + size); } - if (replica.is_disk_replica()) { - // DISK: client_->Get + allocateSlices requires full-object slices, + if (replica.is_disk_replica() || replica.is_dfs_replica()) { + // DISK/DFS: client_->Get + allocateSlices requires full-object slices, // so we must allocate total_size. + const char *replica_type = replica.is_dfs_replica() ? "DFS" : "DISK"; return partial_disk_read( [&](void *tmp_buf) -> tl::expected { std::vector tmp_slices; @@ -3512,7 +3520,7 @@ tl::expected RealClient::execute_ranged_read( auto get_result = client_->Get(key, filtered_qr, tmp_slices); if (!get_result) { LOG(ERROR) - << "DISK Get failed for key: " << key + << replica_type << " Get failed for key: " << key << " with error: " << toString(get_result.error()); return tl::unexpected(get_result.error()); } @@ -3522,7 +3530,7 @@ tl::expected RealClient::execute_ranged_read( } if (!replica.is_memory_replica()) { - LOG(ERROR) << "ranged reads only support memory/disk replicas"; + LOG(ERROR) << "ranged reads only support memory/file-backed replicas"; return tl::unexpected(ErrorCode::INVALID_REPLICA); } @@ -4860,6 +4868,7 @@ RealClient::batch_get_into_internal(const std::vector &keys, std::string key; size_t original_index; QueryResult query_result; + Replica::Descriptor replica; void *dst_buffer; uint64_t total_size; }; @@ -4893,8 +4902,8 @@ RealClient::batch_get_into_internal(const std::vector &keys, continue; } - // Select best replica: prefer local MEMORY, then any MEMORY, - // then LOCAL_DISK, then DISK. + // Select best replica: prefer local MEMORY, any MEMORY, local NOF, + // any NOF, LOCAL_DISK, DFS, then DISK. const auto *best_replica = SelectBestReplica(query_result_values.replicas, local_endpoints); if (!best_replica) { @@ -4929,19 +4938,20 @@ RealClient::batch_get_into_internal(const std::vector &keys, results[i] = static_cast(total_size); continue; } - if (replica.is_disk_replica()) { - // DISK: file I/O (vector_read) cannot write to user GPU buffer. - // Defer — allocate CPU temp buffer, BatchGet, then scatter. + if (replica.is_disk_replica() || replica.is_dfs_replica()) { + // File-backed replicas cannot safely read directly into arbitrary + // user buffers. Use a contiguous temp buffer, then scatter. disk_operations.emplace_back( DiskKeyInfo{.key = key, .original_index = i, .query_result = std::move(query_result_values), + .replica = replica, .dst_buffer = buffers[i], .total_size = total_size}); results[i] = static_cast(total_size); continue; } - // MEMORY: RDMA directly to user buffer. + // MEMORY / NOF: RDMA directly to user buffer. std::vector key_slices; allocateSlices(key_slices, replica, buffers[i]); valid_operations.push_back( @@ -4992,7 +5002,7 @@ RealClient::batch_get_into_internal(const std::vector &keys, } } - // ---- DISK replicas: BatchGet into CPU temp buffers, then scatter ---- + // ---- File-backed replicas: BatchGet into temp buffers, then scatter ---- if (!disk_operations.empty()) { std::vector disk_batch_keys; std::vector disk_batch_qrs; @@ -5003,25 +5013,14 @@ RealClient::batch_get_into_internal(const std::vector &keys, for (size_t di = 0; di < disk_operations.size(); ++di) { auto &op = disk_operations[di]; - // Find the DISK replica. - const Replica::Descriptor *replica_ptr = nullptr; - for (const auto &r : op.query_result.replicas) { - if (r.is_disk_replica()) { - replica_ptr = &r; - break; - } - } - if (!replica_ptr) { - LOG(ERROR) << "No DISK replica found for key: " << op.key; - results[op.original_index] = - tl::unexpected(ErrorCode::INVALID_REPLICA); - continue; - } + const auto &replica = op.replica; + const char *replica_type = + replica.is_dfs_replica() ? "DFS" : "DISK"; auto alloc_result = client_buffer_allocator_->allocate(op.total_size); if (!alloc_result) { - LOG(ERROR) << "Failed to allocate temp buffer for DISK " - << "read, key: " << op.key + LOG(ERROR) << "Failed to allocate temp buffer for " + << replica_type << " read, key: " << op.key << ", size: " << op.total_size; results[op.original_index] = tl::unexpected(ErrorCode::NO_AVAILABLE_HANDLE); @@ -5030,10 +5029,10 @@ RealClient::batch_get_into_internal(const std::vector &keys, auto handle = std::make_unique(std::move(*alloc_result)); std::vector disk_slices; - allocateSlices(disk_slices, *replica_ptr, handle->ptr()); + allocateSlices(disk_slices, replica, handle->ptr()); disk_batch_keys.push_back(op.key); disk_batch_qrs.push_back( - FilterQueryResult(op.query_result, *replica_ptr)); + FilterQueryResult(op.query_result, replica)); disk_batch_slices[op.key] = std::move(disk_slices); disk_batch_indices.push_back(di); disk_temp_handles.emplace(op.key, std::move(handle)); @@ -5047,9 +5046,12 @@ RealClient::batch_get_into_internal(const std::vector &keys, const auto &key = disk_batch_keys[di]; auto &op = disk_operations[disk_batch_indices[di]]; auto handle_it = disk_temp_handles.find(key); + const char *replica_type = + op.replica.is_dfs_replica() ? "DFS" : "DISK"; if (!disk_results[di]) { - LOG(ERROR) << "DISK BatchGet failed for key '" << key - << "': " << toString(disk_results[di].error()); + LOG(ERROR) + << replica_type << " BatchGet failed for key '" << key + << "': " << toString(disk_results[di].error()); results[op.original_index] = tl::unexpected(disk_results[di].error()); continue; @@ -5057,7 +5059,8 @@ RealClient::batch_get_into_internal(const std::vector &keys, if (auto r = scatter_host_to_maybe_device( op.dst_buffer, static_cast(handle_it->second->ptr()), - op.total_size, "DISK read, key: " + key); + op.total_size, + std::string(replica_type) + " read, key: " + key); !r) { results[op.original_index] = tl::make_unexpected(r.error()); } @@ -5923,7 +5926,8 @@ RealClient::batch_get_into_multi_buffers_internal( std::vector buffers; std::vector sizes; uint64_t total_size; - bool is_local_disk; // true=LOCAL_DISK (offload RPC), false=DISK + Replica::Descriptor replica; + bool is_local_disk; // true=LOCAL_DISK (offload RPC), false=DISK/DFS // (BatchGet) }; @@ -5950,9 +5954,9 @@ RealClient::batch_get_into_multi_buffers_internal( results.emplace_back(tl::unexpected(ErrorCode::INVALID_REPLICA)); continue; } - // Select best replica: prefer MEMORY (direct RDMA to GPU), then - // LOCAL_DISK, then DISK. Master may return multiple replicas in any - // order, so always scan rather than blindly taking replicas[0]. + // Select best replica: prefer local MEMORY, any MEMORY, local NOF, + // any NOF, LOCAL_DISK, DFS, then DISK. Master may return multiple + // replicas in any order, so always scan. const auto *best_replica = SelectBestReplica(query_result_values.replicas, local_endpoints); if (!best_replica) { @@ -5978,16 +5982,15 @@ RealClient::batch_get_into_multi_buffers_internal( const auto &buffers = all_buffers[i]; std::vector key_slices; key_slices.reserve(buffers.size()); - if (replica.is_memory_replica()) { - // MEMORY: RDMA from remote memory directly to GPU (GPUDirect). + if (replica.is_memory_replica() || replica.is_nof_replica()) { + // MEMORY / NOF: RDMA directly to the user buffers. for (size_t j = 0; j < buffers.size(); ++j) { key_slices.emplace_back(Slice{buffers[j], sizes[j]}); } } else if (replica.is_local_disk_replica() || - replica.is_disk_replica()) { + replica.is_disk_replica() || replica.is_dfs_replica()) { // LOCAL_DISK: GPU buffers passed directly as scatter-gather slices - // (zero-copy). DISK: file I/O cannot write to GPU memory; temp CPU - // buffer used at read time. + // (zero-copy). DISK/DFS use a contiguous temp buffer at read time. valid_local_disk_ops.emplace( key, DiskKeyInfo{.key = key, @@ -5996,6 +5999,7 @@ RealClient::batch_get_into_multi_buffers_internal( .buffers = all_buffers[i], .sizes = all_sizes[i], .total_size = total_size, + .replica = replica, .is_local_disk = replica.is_local_disk_replica()}); results.emplace_back(static_cast(total_size)); continue; @@ -6019,7 +6023,7 @@ RealClient::batch_get_into_multi_buffers_internal( return results; } - // ---- Memory/Disk replica: existing BatchGet path ---- + // ---- Memory/NOF replica: existing BatchGet path ---- if (!valid_operations.empty()) { std::vector batch_keys; std::vector batch_query_results; @@ -6047,7 +6051,7 @@ RealClient::batch_get_into_multi_buffers_internal( } } - // ---- LOCAL_DISK / DISK replica: disk read paths ---- + // ---- LOCAL_DISK / file-backed replica read paths ---- if (!valid_local_disk_ops.empty()) { // LOCAL_DISK: pass user GPU buffers directly as scatter-gather slices. // vLLM pre-registers all GPU KV-cache memory with TransferEngine via @@ -6062,23 +6066,14 @@ RealClient::batch_get_into_multi_buffers_internal( for (auto &[key, op] : valid_local_disk_ops) { if (!op.is_local_disk) continue; - // Find the correct LOCAL_DISK replica — Master may return - // replicas in any order (e.g. [DISK, LOCAL_DISK]). - const Replica::Descriptor *replica_ptr = nullptr; - for (const auto &r : op.query_result.replicas) { - if (r.is_local_disk_replica()) { - replica_ptr = &r; - break; - } - } - if (!replica_ptr) { + const auto &replica = op.replica; + if (!replica.is_local_disk_replica()) { LOG(ERROR) << "No LOCAL_DISK replica found for key: " << key; results[op.original_index] = tl::make_unexpected(ErrorCode::INVALID_REPLICA); continue; } - const auto &replica = *replica_ptr; std::vector user_slices; user_slices.reserve(op.buffers.size()); size_t slice_total = 0; @@ -6132,10 +6127,9 @@ RealClient::batch_get_into_multi_buffers_internal( } } - // DISK: one batched BatchGet into CPU temp buffers, then scatter. - // (storage_backend::vector_read cannot write directly to GPU memory) + // DISK/DFS: one batched BatchGet into temp buffers, then scatter. { - // Scatter temp CPU buffer -> user multi_buffers (GPU or host). + // Scatter temp buffer -> user multi_buffers (GPU or host). // Returns false and sets results[original_index] on error. auto scatter_to_buffers = [&](const std::string &key, char *src, const DiskKeyInfo &op) -> bool { @@ -6146,8 +6140,12 @@ RealClient::batch_get_into_multi_buffers_internal( std::min(op.sizes[j], static_cast(op.total_size - offset)); void *dst = op.buffers[j]; + const char *replica_type = + op.replica.is_dfs_replica() ? "DFS" : "DISK"; if (auto r = scatter_host_to_maybe_device( - dst, src + offset, sz, "DISK scatter, key: " + key); + dst, src + offset, sz, + std::string(replica_type) + + " scatter, key: " + key); !r) { results[op.original_index] = tl::make_unexpected(r.error()); @@ -6168,39 +6166,26 @@ RealClient::batch_get_into_multi_buffers_internal( for (auto &[key, op] : valid_local_disk_ops) { if (op.is_local_disk) continue; + const auto &replica = op.replica; + const char *replica_type = + replica.is_dfs_replica() ? "DFS" : "DISK"; auto alloc_result = client_buffer_allocator_->allocate(op.total_size); if (!alloc_result) { LOG(ERROR) - << "Failed to allocate temp buffer for DISK " - << "read, key: " << key << ", size: " << op.total_size; + << "Failed to allocate temp buffer for " << replica_type + << " read, key: " << key << ", size: " << op.total_size; results[op.original_index] = tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); continue; } auto handle = std::make_unique(std::move(*alloc_result)); - // Find the correct DISK replica — Master may return - // replicas in any order (e.g. [LOCAL_DISK, DISK]). - const Replica::Descriptor *replica_ptr = nullptr; - for (const auto &r : op.query_result.replicas) { - if (r.is_disk_replica()) { - replica_ptr = &r; - break; - } - } - if (!replica_ptr) { - LOG(ERROR) << "No DISK replica found for key: " << key; - results[op.original_index] = - tl::make_unexpected(ErrorCode::INVALID_REPLICA); - continue; - } - const auto &replica = *replica_ptr; std::vector disk_slices; allocateSlices(disk_slices, replica, handle->ptr()); disk_batch_keys.push_back(key); disk_batch_qrs.push_back( - FilterQueryResult(op.query_result, *replica_ptr)); + FilterQueryResult(op.query_result, replica)); disk_batch_slices[key] = std::move(disk_slices); disk_key_order.push_back(key); temp_handles.emplace(key, std::move(handle)); @@ -6215,9 +6200,12 @@ RealClient::batch_get_into_multi_buffers_internal( const auto &key = disk_key_order[di]; auto &op = valid_local_disk_ops.at(key); auto handle_it = temp_handles.find(key); + const char *replica_type = + op.replica.is_dfs_replica() ? "DFS" : "DISK"; if (!disk_results[di]) { LOG(ERROR) - << "DISK BatchGet failed for key '" << key + << replica_type << " BatchGet failed for key '" + << key << "': " << toString(disk_results[di].error()); results[op.original_index] = tl::make_unexpected(disk_results[di].error()); diff --git a/mooncake-store/src/serialize/serializer.cpp b/mooncake-store/src/serialize/serializer.cpp index ffc9ff8cc1..1392d5b433 100644 --- a/mooncake-store/src/serialize/serializer.cpp +++ b/mooncake-store/src/serialize/serializer.cpp @@ -712,6 +712,24 @@ tl::expected Serializer::serialize( packer.pack(local_data->transport_endpoint); break; } + case ReplicaType::DFS: { + const auto *dfs_data = std::get_if(&replica.data_); + if (!dfs_data) { + return tl::unexpected(SerializationError( + ErrorCode::DESERIALIZE_FAIL, + "serialize_msgpack Replica missing DfsReplicaData")); + } + // Format: [file_path, offset, object_size, aligned_size, shard_idx] + packer.pack_array(5); + packer.pack(dfs_data->descriptor.file_path); + packer.pack(static_cast(dfs_data->descriptor.offset)); + packer.pack( + static_cast(dfs_data->descriptor.object_size)); + packer.pack( + static_cast(dfs_data->descriptor.aligned_size)); + packer.pack(static_cast(dfs_data->descriptor.shard_idx)); + break; + } default: // Unsupported replica type packer.pack(static_cast(255)); @@ -812,6 +830,26 @@ auto Serializer::deserialize(const msgpack::object &obj, client_id, object_size, std::move(transport_endpoint), status); break; } + case static_cast(ReplicaType::DFS): { + const auto &payload = array_items[3]; + if (payload.type != msgpack::type::ARRAY || + payload.via.array.size != 5) { + return tl::unexpected( + SerializationError(ErrorCode::DESERIALIZE_FAIL, + "deserialize_msgpack Replica DFS " + "payload is not valid array[5]")); + } + auto *payload_items = payload.via.array.ptr; + DistributedFSDescriptor descriptor; + descriptor.file_path = payload_items[0].as(); + descriptor.offset = payload_items[1].as(); + descriptor.object_size = payload_items[2].as(); + descriptor.aligned_size = payload_items[3].as(); + descriptor.shard_idx = payload_items[4].as(); + + replica = std::make_shared(std::move(descriptor), status); + break; + } default: return tl::unexpected(SerializationError( ErrorCode::DESERIALIZE_FAIL, diff --git a/mooncake-store/src/storage/distributed/dfs_global_allocator.cpp b/mooncake-store/src/storage/distributed/dfs_global_allocator.cpp new file mode 100644 index 0000000000..3554d704b3 --- /dev/null +++ b/mooncake-store/src/storage/distributed/dfs_global_allocator.cpp @@ -0,0 +1,474 @@ +#include "storage/distributed/dfs_global_allocator.h" + +#include +#include +#include +#include +#include +#include + +#include "storage/distributed/distributed_storage_backend.h" +#include "storage/distributed/fs_adapter.h" +#include "storage/distributed/posix_fs_adapter.h" +#include "utils.h" +#ifdef USE_3FS +#include "storage/distributed/hf3fs_adapter.h" +#endif + +namespace mooncake { + +DfsGlobalAllocator::PendingEviction::~PendingEviction() { + if (owner_ != nullptr) { + owner_->RestorePreparedEviction(std::move(*this)); + } +} + +DfsGlobalAllocator::PendingEviction::PendingEviction( + PendingEviction&& other) noexcept + : owner_(std::exchange(other.owner_, nullptr)), + candidates_(std::move(other.candidates_)), + prepared_(std::move(other.prepared_)) {} + +DfsGlobalAllocator::~DfsGlobalAllocator() { + if (fs_adapter_) fs_adapter_->Shutdown(); +} + +tl::expected DfsGlobalAllocator::Init( + const DistributedStorageConfig& config) { + if (initialized_.load(std::memory_order_acquire)) return {}; + if (!config.ValidateForAllocator()) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + + mount_path_ = config.fsdir; + shard_count_ = config.shard_count; + alignment_ = config.alignment; + shards_.clear(); + shards_.resize(shard_count_); + + eviction_enabled_ = config.eviction_enabled; + eviction_high_watermark_ = config.eviction_high_watermark; + eviction_low_watermark_ = config.eviction_low_watermark; + deferred_free_duration_ = config.deferred_free_duration; + eviction_check_interval_ = config.eviction_check_interval; + + std::error_code ec; + std::filesystem::create_directories(mount_path_, ec); + if (ec) { + LOG(ERROR) << "Failed to create DFS mount path " << mount_path_ << ": " + << ec.message(); + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + + if (config.fs_adapter_type == "posix") { + fs_adapter_ = std::make_unique(); + } else if (config.fs_adapter_type == "hf3fs") { +#ifdef USE_3FS + fs_adapter_ = std::make_unique(); +#else + LOG(ERROR) << "The hf3fs DFS adapter requires Mooncake to be built " + "with the USE_3FS compile-time option " + "(-DUSE_3FS=ON)"; + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); +#endif + } + + auto adapter_init = fs_adapter_->Init(mount_path_); + if (!adapter_init) { + LOG(ERROR) << "Failed to initialize DFS fs adapter " + << config.fs_adapter_type + << " for mount_path=" << mount_path_ + << ", error=" << adapter_init.error(); + return tl::make_unexpected(adapter_init.error()); + } + + for (int i = 0; i < shard_count_; ++i) { + std::string path = mount_path_ + "/dfs_shard_" + + FormatShardIdx(i, shard_count_) + ".data"; + auto prealloc = + fs_adapter_->PreallocateFile(path, config.shard_capacity); + if (!prealloc) { + LOG(ERROR) << "Failed to preallocate DFS shard " << path << ": " + << prealloc.error(); + return tl::make_unexpected(prealloc.error()); + } + + auto shard = std::make_unique(); + shard->capacity = config.shard_capacity; + uint32_t init_cap = static_cast(std::max( + 1, std::min(config.shard_capacity / 4096, 64ULL * 1024))); + uint32_t max_cap = static_cast(std::max( + init_cap, std::min(config.shard_capacity / 1024, + 64ULL * 1024 * 1024))); + shard->allocator = OffsetAllocator::create(0, config.shard_capacity, + init_cap, max_cap); + if (!shard->allocator) { + LOG(ERROR) << "Failed to create offset allocator for DFS shard " + << i; + return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + } + shards_[i] = std::move(shard); + } + + initialized_.store(true, std::memory_order_release); + return {}; +} + +tl::expected DfsGlobalAllocator::Allocate( + const std::string& key, uint64_t size) { + if (!initialized_.load(std::memory_order_acquire)) { + return tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); + } + if (key.empty() || size == 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + + auto key_lock = LockKey(key); + int shard_idx = SelectShard(key); + uint64_t aligned_size = AlignSize(size); + auto& shard = *shards_[shard_idx]; + + std::unique_lock handle_lock(shard.handle_mutex); + ProcessPendingFrees(shard_idx); + + const uint64_t allocation_size = aligned_size + alignment_ - 1; + const uint64_t reserved_bytes = + shard.allocator->normalizedAllocationSize(allocation_size); + auto handle = shard.allocator->allocate(allocation_size); + if (!handle) return tl::make_unexpected(ErrorCode::NO_AVAILABLE_HANDLE); + + uint64_t raw_offset = handle->address(); + uint64_t alloc_offset = AlignSize(raw_offset); + auto alloc_handle = + std::make_shared(std::move(*handle)); + shard.offset_to_handle[alloc_offset] = {key, std::move(alloc_handle), + reserved_bytes}; + handle_lock.unlock(); + + return DistributedFSDescriptor{ + mount_path_ + "/dfs_shard_" + FormatShardIdx(shard_idx, shard_count_) + + ".data", + alloc_offset, + size, + aligned_size, + shard_idx, + }; +} + +void DfsGlobalAllocator::Free(uint64_t offset, uint64_t /*aligned_size*/, + int shard_idx, const std::string& key) { + if (!initialized_.load(std::memory_order_acquire)) return; + if (shard_idx < 0 || shard_idx >= shard_count_) return; + + auto& shard = *shards_[shard_idx]; + std::lock_guard lru_lock(shard.lru_mutex); + std::lock_guard handle_lock(shard.handle_mutex); + + auto lru_it = shard.lru_index.find(key); + if (lru_it != shard.lru_index.end() && lru_it->second->second == offset) { + shard.lru_list.erase(lru_it->second); + shard.lru_index.erase(lru_it); + } + + auto it = shard.offset_to_handle.find(offset); + if (it == shard.offset_to_handle.end()) return; + if (it->second.key != key) return; + + QueuePendingFree( + shard, it->second.handle, it->second.bytes, + std::chrono::steady_clock::now() + deferred_free_duration_); + shard.offset_to_handle.erase(it); +} + +void DfsGlobalAllocator::UpdateAccess(const std::string& key, int shard_idx, + uint64_t offset) { + if (!initialized_.load(std::memory_order_acquire)) return; + if (shard_idx < 0 || shard_idx >= shard_count_) return; + + auto& shard = *shards_[shard_idx]; + std::lock_guard lru_lock(shard.lru_mutex); + { + std::shared_lock handle_lock(shard.handle_mutex); + auto handle_it = shard.offset_to_handle.find(offset); + if (handle_it == shard.offset_to_handle.end() || + handle_it->second.key != key || + handle_it->second.eviction_prepared) { + return; + } + } + auto lru_it = shard.lru_index.find(key); + if (lru_it != shard.lru_index.end()) { + lru_it->second->second = offset; + shard.lru_list.splice(shard.lru_list.begin(), shard.lru_list, + lru_it->second); + } else { + shard.lru_list.push_front({key, offset}); + shard.lru_index[key] = shard.lru_list.begin(); + } +} + +DfsGlobalAllocator::PendingEviction DfsGlobalAllocator::PrepareEviction() { + PendingEviction pending(this); + if (!initialized_.load(std::memory_order_acquire)) return pending; + + for (int i = 0; i < shard_count_; ++i) { + auto& shard = *shards_[i]; + { + std::unique_lock handle_lock(shard.handle_mutex); + CleanupExpiredPendingFrees(shard, std::chrono::steady_clock::now()); + } + PrepareEvictionFromShard(i, pending); + } + return pending; +} + +void DfsGlobalAllocator::CommitPreparedEviction(PendingEviction&& pending) { + if (pending.owner_ != this) return; + + const auto free_at = + std::chrono::steady_clock::now() + deferred_free_duration_; + for (const auto& prepared : pending.prepared_) { + auto& shard = *shards_[prepared.candidate.shard_idx]; + std::lock_guard handle_lock(shard.handle_mutex); + auto it = shard.offset_to_handle.find(prepared.candidate.offset); + if (it == shard.offset_to_handle.end() || + it->second.key != prepared.candidate.key || + it->second.handle != prepared.handle) { + // A concurrent metadata removal may already have called Free(). + continue; + } + + QueuePendingFree(shard, prepared.handle, prepared.bytes, free_at); + shard.offset_to_handle.erase(it); + } + + pending.prepared_.clear(); + pending.candidates_.clear(); + pending.owner_ = nullptr; +} + +void DfsGlobalAllocator::RestorePreparedEviction(PendingEviction&& pending) { + if (pending.owner_ != this) return; + + // Candidates were removed oldest-first. Restore in reverse order so the + // original relative LRU order is preserved. + for (auto prepared_it = pending.prepared_.rbegin(); + prepared_it != pending.prepared_.rend(); ++prepared_it) { + const auto& prepared = *prepared_it; + auto& shard = *shards_[prepared.candidate.shard_idx]; + std::lock_guard lru_lock(shard.lru_mutex); + std::lock_guard handle_lock(shard.handle_mutex); + + auto handle_it = shard.offset_to_handle.find(prepared.candidate.offset); + if (handle_it == shard.offset_to_handle.end() || + handle_it->second.key != prepared.candidate.key || + handle_it->second.handle != prepared.handle) { + // Free() or a replacement already retired this allocation. + continue; + } + + handle_it->second.eviction_prepared = false; + if (shard.lru_index.find(prepared.candidate.key) == + shard.lru_index.end()) { + shard.lru_list.push_back( + {prepared.candidate.key, prepared.candidate.offset}); + shard.lru_index[prepared.candidate.key] = + std::prev(shard.lru_list.end()); + } + } + + pending.prepared_.clear(); + pending.candidates_.clear(); + pending.owner_ = nullptr; +} + +void DfsGlobalAllocator::ResolvePreparedEviction( + PendingEviction&& pending, const std::vector& accepted) { + if (pending.owner_ != this) return; + if (accepted.size() != pending.prepared_.size()) { + LOG(ERROR) << "DFS eviction decision count " << accepted.size() + << " does not match prepared candidate count " + << pending.prepared_.size(); + RestorePreparedEviction(std::move(pending)); + return; + } + + const auto free_at = + std::chrono::steady_clock::now() + deferred_free_duration_; + for (size_t i = 0; i < pending.prepared_.size(); ++i) { + if (!accepted[i]) continue; + + const auto& prepared = pending.prepared_[i]; + auto& shard = *shards_[prepared.candidate.shard_idx]; + std::lock_guard handle_lock(shard.handle_mutex); + auto handle_it = shard.offset_to_handle.find(prepared.candidate.offset); + if (handle_it == shard.offset_to_handle.end() || + handle_it->second.key != prepared.candidate.key || + handle_it->second.handle != prepared.handle) { + // Free() or a replacement already retired this allocation. + continue; + } + + QueuePendingFree(shard, prepared.handle, prepared.bytes, free_at); + shard.offset_to_handle.erase(handle_it); + } + + // Candidates were removed oldest-first. Restore rejected entries in + // reverse order so their relative LRU order is preserved. + for (size_t i = pending.prepared_.size(); i > 0; --i) { + if (accepted[i - 1]) continue; + + const auto& prepared = pending.prepared_[i - 1]; + auto& shard = *shards_[prepared.candidate.shard_idx]; + std::lock_guard lru_lock(shard.lru_mutex); + std::lock_guard handle_lock(shard.handle_mutex); + + auto handle_it = shard.offset_to_handle.find(prepared.candidate.offset); + if (handle_it == shard.offset_to_handle.end() || + handle_it->second.key != prepared.candidate.key || + handle_it->second.handle != prepared.handle) { + continue; + } + + handle_it->second.eviction_prepared = false; + if (shard.lru_index.find(prepared.candidate.key) == + shard.lru_index.end()) { + shard.lru_list.push_back( + {prepared.candidate.key, prepared.candidate.offset}); + shard.lru_index[prepared.candidate.key] = + std::prev(shard.lru_list.end()); + } + } + + pending.prepared_.clear(); + pending.candidates_.clear(); + pending.owner_ = nullptr; +} + +std::string DfsGlobalAllocator::FormatShardIdx(int idx, int shard_count) { + int width = static_cast(std::max( + 2, std::to_string(std::max(0, shard_count - 1)).size())); + std::ostringstream oss; + oss << std::setw(width) << std::setfill('0') << idx; + return oss.str(); +} + +void DfsGlobalAllocator::ProcessPendingFrees(int shard_idx) { + auto& shard = *shards_[shard_idx]; + CleanupExpiredPendingFrees(shard, std::chrono::steady_clock::now()); +} + +void DfsGlobalAllocator::QueuePendingFree( + ShardState& shard, const std::shared_ptr& handle, + uint64_t bytes, std::chrono::steady_clock::time_point when) { + if (!handle) return; + std::lock_guard pending_lock(shard.pending_mutex); + if (bytes == 0) bytes = handle->size(); + shard.pending_free.push_back({handle, bytes, when}); + shard.pending_free_bytes += bytes; +} + +void DfsGlobalAllocator::CleanupExpiredPendingFrees( + ShardState& shard, std::chrono::steady_clock::time_point now) { + std::lock_guard pending_lock(shard.pending_mutex); + while (!shard.pending_free.empty() && + shard.pending_free.front().when <= now) { + const uint64_t bytes = shard.pending_free.front().bytes; + shard.pending_free.pop_front(); + if (bytes > shard.pending_free_bytes) { + shard.pending_free_bytes = 0; + } else { + shard.pending_free_bytes -= bytes; + } + } +} + +double DfsGlobalAllocator::EffectiveUsage(ShardState& shard) { + uint64_t physical_free = 0; + { + std::shared_lock lock(shard.handle_mutex); + auto report = shard.allocator->storageReport(); + physical_free = report.totalFreeSpace; + } + + uint64_t pending_free_bytes = 0; + { + std::lock_guard pending_lock(shard.pending_mutex); + pending_free_bytes = shard.pending_free_bytes; + } + + const uint64_t capped_physical_free = + std::min(physical_free, shard.capacity); + const uint64_t remaining_capacity = shard.capacity - capped_physical_free; + const uint64_t effective_free = + capped_physical_free + std::min(pending_free_bytes, remaining_capacity); + if (shard.capacity == 0) return 0.0; + return 1.0 - static_cast(effective_free) / + static_cast(shard.capacity); +} + +void DfsGlobalAllocator::PrepareEvictionFromShard(int shard_idx, + PendingEviction& pending) { + auto& shard = *shards_[shard_idx]; + + const double usage = EffectiveUsage(shard); + uint64_t prepared_bytes = 0; + std::lock_guard lru_lock(shard.lru_mutex); + std::lock_guard handle_lock(shard.handle_mutex); + + if (usage >= eviction_high_watermark_) { + shard.eviction_active = true; + } + if (!shard.eviction_active) return; + if (usage < eviction_low_watermark_) { + shard.eviction_active = false; + return; + } + + while (true) { + if (shard.lru_list.empty()) break; + auto lru_it = std::prev(shard.lru_list.end()); + const std::string evict_key = lru_it->first; + const uint64_t evict_offset = lru_it->second; + + auto handle_it = shard.offset_to_handle.find(evict_offset); + if (handle_it == shard.offset_to_handle.end() || + handle_it->second.key != evict_key || + handle_it->second.eviction_prepared) { + shard.lru_list.erase(lru_it); + shard.lru_index.erase(evict_key); + continue; + } + + EvictionCandidate candidate{evict_key, shard_idx, evict_offset}; + PendingEviction::PreparedAllocation prepared{ + candidate, handle_it->second.handle, handle_it->second.bytes}; + pending.prepared_.push_back(std::move(prepared)); + try { + pending.candidates_.push_back(std::move(candidate)); + } catch (...) { + pending.prepared_.pop_back(); + throw; + } + + handle_it->second.eviction_prepared = true; + prepared_bytes += handle_it->second.bytes; + shard.lru_list.erase(lru_it); + shard.lru_index.erase(evict_key); + + const double projected_usage = + usage - static_cast(prepared_bytes) / + static_cast(shard.capacity); + if (projected_usage < eviction_low_watermark_) break; + } +} + +int DfsGlobalAllocator::SelectShard(const std::string& key) const { + return std::hash{}(key) % shard_count_; +} + +uint64_t DfsGlobalAllocator::AlignSize(uint64_t size) const { + return (size + alignment_ - 1) & ~(alignment_ - 1); +} + +} // namespace mooncake diff --git a/mooncake-store/src/storage/distributed/distributed_storage_backend.cpp b/mooncake-store/src/storage/distributed/distributed_storage_backend.cpp index 56483d05f6..7eb7b4d7ee 100644 --- a/mooncake-store/src/storage/distributed/distributed_storage_backend.cpp +++ b/mooncake-store/src/storage/distributed/distributed_storage_backend.cpp @@ -1,18 +1,39 @@ #include "storage/distributed/distributed_storage_backend.h" -#include - #include -#include -#include #include +#include +#include #include "environ.h" +#include "storage/distributed/dfs_global_allocator.h" +#include "types.h" #include "utils.h" namespace mooncake { -// === DistributedStorageConfig === +namespace { + +bool IsDfsDescriptorRangeValid(const DistributedFSDescriptor& desc, + const DistributedStorageConfig& config) { + if (config.alignment == 0 || desc.object_size == 0 || + desc.aligned_size < desc.object_size || + desc.offset % config.alignment != 0 || + desc.aligned_size % config.alignment != 0) { + return false; + } + if (desc.offset > config.shard_capacity || + desc.aligned_size > config.shard_capacity - desc.offset) { + return false; + } + + constexpr uint64_t kMaxFileOffset = + static_cast(std::numeric_limits::max()); + return desc.offset <= kMaxFileOffset && + desc.aligned_size <= kMaxFileOffset - desc.offset; +} + +} // namespace bool DistributedStorageConfig::Validate() const { if (fsdir.empty()) { @@ -25,17 +46,57 @@ bool DistributedStorageConfig::Validate() const { << fsdir; return false; } - if (fs_adapter_type.empty()) { - LOG(ERROR) << "DistributedStorageConfig: fs_adapter_type is empty"; - return false; - } - if (fs_adapter_type != "hf3fs") { + if (fs_adapter_type != "hf3fs" && fs_adapter_type != "posix") { LOG(ERROR) << "DistributedStorageConfig: unsupported fs_adapter_type: " << fs_adapter_type; return false; } - if (hash_bucket_count <= 0) { - LOG(ERROR) << "DistributedStorageConfig: hash_bucket_count must > 0"; + if (shard_count <= 0) { + LOG(ERROR) << "DistributedStorageConfig: shard_count must > 0"; + return false; + } + if (shard_capacity == 0) { + LOG(ERROR) << "DistributedStorageConfig: shard_capacity must > 0"; + return false; + } + if (alignment == 0 || (alignment & (alignment - 1)) != 0) { + LOG(ERROR) << "DistributedStorageConfig: alignment must be power of 2"; + return false; + } + if (shard_capacity % alignment != 0) { + LOG(ERROR) << "DistributedStorageConfig: shard_capacity must align"; + return false; + } + if (!single_tenant) { + LOG(ERROR) << "DistributedStorageConfig: Currently, DFS requires " + "single_tenant=true"; + return false; + } + return true; +} + +bool DistributedStorageConfig::ValidateForAllocator() const { + if (!Validate()) return false; + + if (eviction_low_watermark < 0.0 || eviction_low_watermark > 1.0 || + eviction_high_watermark < 0.0 || eviction_high_watermark > 1.0 || + eviction_low_watermark >= eviction_high_watermark) { + LOG(ERROR) << "DistributedStorageConfig: eviction watermarks must " + "satisfy 0 <= low < high <= 1, low=" + << eviction_low_watermark + << ", high=" << eviction_high_watermark; + return false; + } + if (deferred_free_duration.count() < 0) { + LOG(ERROR) << "DistributedStorageConfig: deferred_free_duration must " + "be non-negative, seconds=" + << deferred_free_duration.count(); + return false; + } + if (eviction_enabled && eviction_check_interval.count() <= 0) { + LOG(ERROR) << "DistributedStorageConfig: eviction_check_interval must " + "be positive when eviction is enabled, seconds=" + << eviction_check_interval.count(); return false; } return true; @@ -43,21 +104,56 @@ bool DistributedStorageConfig::Validate() const { DistributedStorageConfig DistributedStorageConfig::FromEnvironment() { DistributedStorageConfig config; - config.fsdir = - Environ::GetString("MOONCAKE_DISTRIBUTED_ROOT_DIR", config.fsdir); + config.fsdir = Environ::GetString( + "MOONCAKE_DFS_ROOT_DIR", + Environ::GetString("MOONCAKE_DISTRIBUTED_ROOT_DIR", config.fsdir)); if (!std::filesystem::path(config.fsdir).is_absolute()) { config.fsdir = std::filesystem::absolute(config.fsdir).string(); } - config.fs_adapter_type = Environ::GetString("MOONCAKE_DISTRIBUTED_FS_TYPE", - config.fs_adapter_type); + config.fs_adapter_type = + Environ::GetString("MOONCAKE_DFS_FS_ADAPTER", + Environ::GetString("MOONCAKE_DISTRIBUTED_FS_TYPE", + config.fs_adapter_type)); config.enable_health_check = Environ::GetBool("MOONCAKE_DISTRIBUTED_HEALTH_CHECK", false); - config.hash_bucket_count = - Environ::GetInt("MOONCAKE_DISTRIBUTED_HASH_BUCKET_COUNT", 256); + config.shard_count = + Environ::GetInt("MOONCAKE_DFS_SHARD_COUNT", config.shard_count); + config.shard_capacity = Environ::GetUInt64("MOONCAKE_DFS_SHARD_CAPACITY", + config.shard_capacity); + config.alignment = + Environ::GetUInt64("MOONCAKE_DFS_ALIGNMENT", config.alignment); + config.single_tenant = + Environ::GetBool("MOONCAKE_DFS_SINGLE_TENANT", config.single_tenant); + config.eviction_enabled = Environ::GetBool("MOONCAKE_DFS_EVICTION_ENABLED", + config.eviction_enabled); + config.eviction_high_watermark = Environ::GetDouble( + "MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", config.eviction_high_watermark); + config.eviction_low_watermark = Environ::GetDouble( + "MOONCAKE_DFS_EVICTION_LOW_WATERMARK", config.eviction_low_watermark); + config.deferred_free_duration = std::chrono::seconds(Environ::GetInt( + "MOONCAKE_DFS_DEFERRED_FREE_SECONDS", + static_cast(config.deferred_free_duration.count()))); + config.eviction_check_interval = std::chrono::seconds(Environ::GetInt( + "MOONCAKE_DFS_EVICTION_CHECK_INTERVAL", + static_cast(config.eviction_check_interval.count()))); return config; } -// === DistributedStorageBackend === +std::string DistributedStorageConfig::FormatStr() const { + std::ostringstream oss; + oss << "fsdir=" << fsdir << ", fs_adapter_type=" << fs_adapter_type + << ", enable_health_check=" << enable_health_check + << ", shard_count=" << shard_count + << ", shard_capacity=" << shard_capacity << ", alignment=" << alignment + << ", single_tenant=" << single_tenant + << ", eviction_enabled=" << eviction_enabled + << ", eviction_high_watermark=" << eviction_high_watermark + << ", eviction_low_watermark=" << eviction_low_watermark + << ", deferred_free_seconds=" << deferred_free_duration.count() + << ", eviction_check_interval_seconds=" + << eviction_check_interval.count(); + return oss.str(); +} DistributedStorageBackend::DistributedStorageBackend( const FileStorageConfig& file_storage_config, @@ -66,69 +162,50 @@ DistributedStorageBackend::DistributedStorageBackend( : StorageBackendInterface(file_storage_config), fs_adapter_(std::move(fs_adapter)), distributed_config_(distributed_config), - root_dir_(distributed_config.fsdir), - hash_bucket_count_(distributed_config.hash_bucket_count) {} + root_dir_(distributed_config.fsdir) {} + +DistributedStorageBackend::~DistributedStorageBackend() { + for (auto& shard : shard_files_) { + if (shard && shard->fd >= 0 && fs_adapter_) { + fs_adapter_->CloseFile(shard->fd); + shard->fd = -1; + } + } + if (fs_adapter_) fs_adapter_->Shutdown(); +} tl::expected DistributedStorageBackend::Init() { if (initialized_) { LOG(WARNING) << "DistributedStorageBackend is already initialized"; return {}; } - - auto init_result = fs_adapter_->Init(root_dir_); - if (!init_result) return init_result; - - // Ensure root directory exists before health check std::error_code ec; std::filesystem::create_directories(root_dir_, ec); if (ec) { - LOG(ERROR) << "Failed to create root directory " << root_dir_ << ": " - << ec.message(); + LOG(ERROR) << "Failed to create DFS root directory " << root_dir_ + << ": " << ec.message(); return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); } - if (distributed_config_.enable_health_check) { - std::string probe_path = - fmt::format("{}/.mooncake_health_probe_{}", root_dir_, - UuidToString(generate_uuid())); - std::string probe_data = "health_check"; - auto write_result = fs_adapter_->WriteFile( - probe_path, - std::span(probe_data.data(), probe_data.size())); - if (!write_result) { - LOG(ERROR) << "DFS health check failed (write): " - << static_cast(write_result.error()); - return tl::make_unexpected(write_result.error()); - } - - std::vector read_buf(probe_data.size()); - auto read_result = - fs_adapter_->ReadFile(probe_path, read_buf.data(), read_buf.size()); - if (!read_result || *read_result != probe_data.size() || - std::string(read_buf.data(), read_buf.size()) != probe_data) { - LOG(ERROR) << "DFS health check failed (read back mismatch)"; - auto del_err = fs_adapter_->DeleteFile(probe_path); - if (!del_err) { - LOG(WARNING) << "Failed to delete health-check probe: " - << static_cast(del_err.error()); - } - return tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE); - } - - fs_adapter_->DeleteFile(probe_path); - LOG(INFO) << "DFS health check passed, adapter=" - << fs_adapter_->GetName(); - } + auto init_result = fs_adapter_->Init(root_dir_); + if (!init_result) return init_result; - // Ensure hash bucket directories exist - for (int i = 0; i < hash_bucket_count_; ++i) { - std::string bucket_dir = fmt::format("{}/{:02x}", root_dir_, i); - std::filesystem::create_directories(bucket_dir, ec); - if (ec) { - LOG(ERROR) << "Failed to create bucket directory " << bucket_dir - << ": " << ec.message(); - return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + shard_files_.reserve(distributed_config_.shard_count); + for (int i = 0; i < distributed_config_.shard_count; ++i) { + std::string path = root_dir_ + "/dfs_shard_" + + DfsGlobalAllocator::FormatShardIdx( + i, distributed_config_.shard_count) + + ".data"; + auto fd_result = fs_adapter_->OpenFile(path); + if (!fd_result) { + LOG(ERROR) << "Failed to open DFS shard " << path << ": " + << fd_result.error(); + return tl::make_unexpected(fd_result.error()); } + auto shard = std::make_unique(); + shard->path = std::move(path); + shard->fd = *fd_result; + shard_files_.push_back(std::move(shard)); } initialized_ = true; @@ -136,185 +213,222 @@ tl::expected DistributedStorageBackend::Init() { } tl::expected DistributedStorageBackend::BatchOffload( - const std::unordered_map>& batch_object, + const std::unordered_map>& /*batch_object*/, std::function& keys, std::vector& metadatas)> - complete_handler, - EvictionHandler eviction_handler) { + /*complete_handler*/, + EvictionHandler /*eviction_handler*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); +} + +std::vector> +DistributedStorageBackend::BatchWrite( + const std::vector& requests) { + std::vector> results; + results.reserve(requests.size()); + if (!initialized_) { LOG(ERROR) << "DistributedStorageBackend is not initialized"; - return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + results.assign(requests.size(), + tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE)); + return results; } - if (eviction_handler) { - LOG_FIRST_N(WARNING, 1) - << "DistributedStorageBackend does not support eviction, " - "eviction_handler ignored"; - } - - std::vector success_keys; - std::vector success_metas; + for (const auto& request : requests) { + const auto& desc = request.descriptor; + if (desc.shard_idx < 0 || + desc.shard_idx >= static_cast(shard_files_.size())) { + LOG(ERROR) << "Invalid DFS shard_idx " << desc.shard_idx + << " for key " << request.key; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; + } - for (const auto& [key, slices] : batch_object) { - auto path = GetObjectPath(key); + auto& shard = *shard_files_[desc.shard_idx]; + if (desc.file_path != shard.path) { + LOG(ERROR) << "DFS path mismatch for key " << request.key + << ", descriptor=" << desc.file_path + << ", configured=" << shard.path; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; + } + if (!IsDfsDescriptorRangeValid(desc, distributed_config_)) { + LOG(ERROR) << "Invalid DFS descriptor range for key " << request.key + << ", offset=" << desc.offset + << ", object_size=" << desc.object_size + << ", aligned_size=" << desc.aligned_size + << ", shard_capacity=" + << distributed_config_.shard_capacity; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; + } std::vector iovs; - for (const auto& slice : slices) { + iovs.reserve(request.slices.size()); + uint64_t total_size = 0; + bool invalid = false; + for (const auto& slice : request.slices) { + if ((!slice.ptr && slice.size > 0) || + slice.size > + std::numeric_limits::max() - total_size) { + invalid = true; + break; + } + total_size += slice.size; iovs.push_back({slice.ptr, slice.size}); } - - auto result = - fs_adapter_->VectorWriteFile(path, iovs.data(), iovs.size(), 0); - if (!result) { - LOG(WARNING) << "Failed to offload key " << key << ": " - << static_cast(result.error()); + if (invalid || total_size != desc.object_size) { + LOG(WARNING) << "Invalid DFS write request for key " << request.key + << ", expected=" << desc.object_size + << ", actual=" << total_size; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); continue; } - success_keys.push_back(key); - StorageObjectMetadata meta{-1, 0, static_cast(key.size()), - static_cast(*result), ""}; - success_metas.push_back(meta); - } - - if (!success_keys.empty()) { - auto err = complete_handler(success_keys, success_metas); - if (err != ErrorCode::OK) { - return tl::make_unexpected(err); + std::lock_guard lock(shard.mutex); + auto write_result = + fs_adapter_->WriteAt(shard.fd, iovs.data(), iovs.size(), + static_cast(desc.offset)); + if (!write_result) { + LOG(WARNING) << "DFS write failed for key " << request.key + << ", error=" << write_result.error(); + results.emplace_back(tl::make_unexpected(write_result.error())); + continue; } + if (*write_result != total_size) { + LOG(WARNING) << "DFS short write for key " << request.key + << ", expected=" << total_size + << ", actual=" << *write_result; + results.emplace_back( + tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL)); + continue; + } + results.emplace_back(); } - - return static_cast(success_keys.size()); + return results; } -tl::expected DistributedStorageBackend::BatchLoad( - std::unordered_map& batched_slices) { +std::vector> DistributedStorageBackend::BatchRead( + const std::vector& requests) { + std::vector> results; + results.reserve(requests.size()); + if (!initialized_) { LOG(ERROR) << "DistributedStorageBackend is not initialized"; - return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); + results.assign(requests.size(), + tl::make_unexpected(ErrorCode::DFS_SERVICE_UNAVAILABLE)); + return results; } - for (auto& [key, slice] : batched_slices) { - auto path = GetObjectPath(key); + for (const auto& request : requests) { + const auto& desc = request.descriptor; + if (desc.shard_idx < 0 || + desc.shard_idx >= static_cast(shard_files_.size())) { + LOG(ERROR) << "Invalid DFS shard_idx " << desc.shard_idx + << " for key " << request.key; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; + } - auto result = fs_adapter_->ReadFile(path, slice.ptr, slice.size); - if (!result) { - return tl::make_unexpected(result.error()); + auto& shard = *shard_files_[desc.shard_idx]; + if (desc.file_path != shard.path) { + LOG(ERROR) << "DFS path mismatch for key " << request.key + << ", descriptor=" << desc.file_path + << ", configured=" << shard.path; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; } - if (*result != slice.size) { - return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + if (!IsDfsDescriptorRangeValid(desc, distributed_config_)) { + LOG(ERROR) << "Invalid DFS descriptor range for key " << request.key + << ", offset=" << desc.offset + << ", object_size=" << desc.object_size + << ", aligned_size=" << desc.aligned_size + << ", shard_capacity=" + << distributed_config_.shard_capacity; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; + } + if (desc.object_size > std::numeric_limits::max() || + request.slices.size() > + static_cast(std::numeric_limits::max())) { + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; } - } - return {}; -} - -tl::expected DistributedStorageBackend::IsExist( - const std::string& key) { - if (!initialized_) { - LOG(ERROR) << "DistributedStorageBackend is not initialized"; - return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); - } - - auto path = GetObjectPath(key); - return fs_adapter_->FileExists(path); -} - -tl::expected DistributedStorageBackend::IsEnableOffloading() { - return true; -} - -tl::expected DistributedStorageBackend::ScanMeta( - const std::function< - ErrorCode(const std::vector& keys, - std::vector& metadatas)>& handler) { - if (!initialized_) { - LOG(ERROR) << "DistributedStorageBackend is not initialized"; - return tl::make_unexpected(ErrorCode::INTERNAL_ERROR); - } - - std::vector batch_keys; - std::vector batch_metas; - const size_t batch_limit = static_cast(std::max( - 1, file_storage_config_.scanmeta_iterator_keys_limit)); - for (int i = 0; i < hash_bucket_count_; ++i) { - std::string bucket_dir = fmt::format("{}/{:02x}", root_dir_, i); - auto file_infos = fs_adapter_->ListFilesWithInfo(bucket_dir); - if (!file_infos) { - if (file_infos.error() == ErrorCode::FILE_NOT_FOUND) { + std::vector iovs; + iovs.reserve(request.slices.size()); + size_t remaining = static_cast(desc.object_size); + bool invalid = false; + for (const auto& slice : request.slices) { + if (!slice.ptr && slice.size > 0) { + invalid = true; + break; + } + if (remaining == 0 || slice.size == 0) { continue; } - LOG(ERROR) << "Failed to list files in bucket " << bucket_dir - << ": " << static_cast(file_infos.error()); - return tl::make_unexpected(file_infos.error()); + const size_t read_size = std::min(slice.size, remaining); + iovs.push_back({slice.ptr, read_size}); + remaining -= read_size; + } + if (invalid || remaining != 0) { + LOG(WARNING) << "Invalid DFS read request for key " << request.key + << ", expected capacity at least=" << desc.object_size; + results.emplace_back( + tl::make_unexpected(ErrorCode::INVALID_PARAMS)); + continue; } - for (const auto& info : *file_infos) { - std::string key = UnescapeFilename(info.name); - batch_keys.push_back(key); - StorageObjectMetadata meta{-1, 0, static_cast(key.size()), - static_cast(info.size), ""}; - batch_metas.push_back(meta); - - if (batch_keys.size() >= batch_limit) { - auto err = handler(batch_keys, batch_metas); - if (err != ErrorCode::OK) return tl::make_unexpected(err); - batch_keys.clear(); - batch_metas.clear(); - } + std::lock_guard lock(shard.mutex); + auto read_result = fs_adapter_->ReadAt( + shard.fd, iovs.data(), static_cast(iovs.size()), + static_cast(desc.offset)); + if (!read_result) { + LOG(WARNING) << "DFS read failed for key " << request.key + << ", error=" << read_result.error(); + results.emplace_back(tl::make_unexpected(read_result.error())); + continue; } + if (*read_result != desc.object_size) { + LOG(WARNING) << "DFS short read for key " << request.key + << ", expected=" << desc.object_size + << ", actual=" << *read_result; + results.emplace_back( + tl::make_unexpected(ErrorCode::FILE_READ_FAIL)); + continue; + } + results.emplace_back(); } - if (!batch_keys.empty()) { - auto err = handler(batch_keys, batch_metas); - if (err != ErrorCode::OK) return tl::make_unexpected(err); - } - return {}; + return results; } -// === Key -> Path mapping === +tl::expected DistributedStorageBackend::BatchLoad( + std::unordered_map& /*batched_slices*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); +} -std::string DistributedStorageBackend::GetObjectPath( - const std::string& key) const { - uint64_t hash = XXH64(key.data(), key.size(), 0); - std::string bucket = fmt::format("{:02x}", hash % hash_bucket_count_); - std::string safe_key = EscapeFilename(key); - return (std::filesystem::path(root_dir_) / bucket / safe_key).string(); +tl::expected DistributedStorageBackend::IsExist( + const std::string& /*key*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); } -std::string DistributedStorageBackend::EscapeFilename(const std::string& key) { - std::string result; - result.reserve(key.size() + 16); - for (unsigned char c : key) { - if (c == '@' || c == ':' || c == '/' || c == '\\' || c == '%' || - c < 0x20 || c > 0x7e) { - result += fmt::format("%{:02x}", static_cast(c)); - } else { - result += static_cast(c); - } - } - return result; +tl::expected DistributedStorageBackend::IsEnableOffloading() { + return false; } -std::string DistributedStorageBackend::UnescapeFilename( - const std::string& name) { - auto is_hex = [](char c) { - return (c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || - (c >= 'A' && c <= 'F'); - }; - std::string result; - result.reserve(name.size()); - for (size_t i = 0; i < name.size(); ++i) { - if (name[i] == '%' && i + 2 < name.size() && is_hex(name[i + 1]) && - is_hex(name[i + 2])) { - char hex[3] = {name[i + 1], name[i + 2], 0}; - unsigned long val = strtoul(hex, nullptr, 16); - result += static_cast(val); - i += 2; - } else { - result += name[i]; - } - } - return result; +tl::expected DistributedStorageBackend::ScanMeta( + const std::function& keys, + std::vector& metadatas)>& /*handler*/) { + return tl::make_unexpected(ErrorCode::NOT_SUPPORTED); } } // namespace mooncake diff --git a/mooncake-store/src/storage/distributed/hf3fs_adapter.cpp b/mooncake-store/src/storage/distributed/hf3fs_adapter.cpp index a4ccf3b522..4fe853c676 100644 --- a/mooncake-store/src/storage/distributed/hf3fs_adapter.cpp +++ b/mooncake-store/src/storage/distributed/hf3fs_adapter.cpp @@ -2,12 +2,16 @@ #include #include +#include #include +#include #include #include #include +#include + #include "hf3fs/hf3fs.h" namespace mooncake { @@ -21,7 +25,14 @@ tl::expected Hf3fsAdapter::Init( const std::string& mount_path) { resource_manager_ = std::make_unique(); Hf3fsConfig config{}; - config.mount_root = mount_path; + char hf3fs_mount_point[PATH_MAX] = {}; + int ret = hf3fs_extract_mount_point( + hf3fs_mount_point, sizeof(hf3fs_mount_point), mount_path.c_str()); + if (ret > 0 && ret <= static_cast(sizeof(hf3fs_mount_point))) { + config.mount_root = hf3fs_mount_point; + } else { + config.mount_root = mount_path; + } resource_manager_->setDefaultParams(config); return {}; } @@ -379,4 +390,185 @@ tl::expected, ErrorCode> Hf3fsAdapter::ListFiles( return result; } +tl::expected Hf3fsAdapter::OpenFile(const std::string& path) { + int fd = open(path.c_str(), O_RDWR | O_CREAT | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + + if (hf3fs_reg_fd(fd, 0) > 0) { + close(fd); + return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + } + return fd; +} + +tl::expected Hf3fsAdapter::CloseFile(int fd) { + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_INVALID_HANDLE); + hf3fs_dereg_fd(fd); + if (close(fd) != 0) { + return tl::make_unexpected(ErrorCode::FILE_INVALID_HANDLE); + } + return {}; +} + +tl::expected Hf3fsAdapter::PreallocateFile( + const std::string& path, uint64_t size) { + int fd = open(path.c_str(), O_RDWR | O_CREAT | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + + int rc = fallocate(fd, 0, 0, static_cast(size)); + if (rc != 0) { + rc = ftruncate(fd, static_cast(size)); + } + int saved_errno = errno; + close(fd); + + if (rc != 0) { + errno = saved_errno; + LOG(ERROR) << "Failed to preallocate DFS file " << path + << ", size=" << size << ", error=" << strerror(errno); + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + return {}; +} + +tl::expected Hf3fsAdapter::WriteAt(int fd, const iovec* iov, + int iovcnt, + int64_t offset) { + if (fd < 0 || offset < 0 || iovcnt < 0 || (iovcnt > 0 && iov == nullptr)) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + for (int i = 0; i < iovcnt; ++i) { + if (!iov[i].iov_base && iov[i].iov_len > 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + } + + auto* resource = resource_manager_->getThreadResource(); + if (!resource || !resource->initialized) { + return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + } + + size_t total_length = 0; + for (int i = 0; i < iovcnt; ++i) total_length += iov[i].iov_len; + if (total_length == 0) return size_t{0}; + + auto& threefs_iov = resource->iov_; + auto& ior_write = resource->ior_write_; + size_t total_written = 0; + size_t remaining = total_length; + int iov_idx = 0; + size_t iov_off = 0; + off_t current_offset = static_cast(offset); + + while (remaining > 0) { + size_t chunk = std::min(remaining, resource->config_.iov_size); + size_t copied = 0; + char* dest = reinterpret_cast(threefs_iov.base); + while (copied < chunk && iov_idx < iovcnt) { + if (iov_off >= iov[iov_idx].iov_len) { + ++iov_idx; + iov_off = 0; + continue; + } + size_t n = std::min(chunk - copied, iov[iov_idx].iov_len - iov_off); + memcpy(dest + copied, + static_cast(iov[iov_idx].iov_base) + iov_off, n); + copied += n; + iov_off += n; + } + if (copied == 0) break; + + int ret = + hf3fs_prep_io(&ior_write, &threefs_iov, false, threefs_iov.base, fd, + current_offset, copied, nullptr); + if (ret < 0) break; + ret = hf3fs_submit_ios(&ior_write); + if (ret < 0) break; + struct hf3fs_cqe cqe; + ret = hf3fs_wait_for_ios(&ior_write, &cqe, 1, 1, nullptr); + if (ret < 0 || cqe.result < 0) break; + + size_t bytes_written = cqe.result; + if (bytes_written == 0) break; + total_written += bytes_written; + current_offset += bytes_written; + remaining -= bytes_written; + if (bytes_written < copied) break; + } + + if (total_written != total_length) { + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + return total_written; +} + +tl::expected Hf3fsAdapter::ReadAt(int fd, iovec* iov, + int iovcnt, + int64_t offset) { + if (fd < 0 || offset < 0 || iovcnt < 0 || (iovcnt > 0 && iov == nullptr)) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + for (int i = 0; i < iovcnt; ++i) { + if (!iov[i].iov_base && iov[i].iov_len > 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + } + + auto* resource = resource_manager_->getThreadResource(); + if (!resource || !resource->initialized) { + return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + } + + size_t total_length = 0; + for (int i = 0; i < iovcnt; ++i) total_length += iov[i].iov_len; + if (total_length == 0) return size_t{0}; + + auto& threefs_iov = resource->iov_; + auto& ior_read = resource->ior_read_; + size_t total_read = 0; + size_t remaining = total_length; + int iov_idx = 0; + size_t iov_off = 0; + off_t current_offset = static_cast(offset); + + while (remaining > 0) { + size_t chunk = std::min(remaining, resource->config_.iov_size); + int ret = hf3fs_prep_io(&ior_read, &threefs_iov, true, threefs_iov.base, + fd, current_offset, chunk, nullptr); + if (ret < 0) break; + ret = hf3fs_submit_ios(&ior_read); + if (ret < 0) break; + struct hf3fs_cqe cqe; + ret = hf3fs_wait_for_ios(&ior_read, &cqe, 1, 1, nullptr); + if (ret < 0 || cqe.result < 0) break; + + size_t bytes_read = cqe.result; + if (bytes_read == 0) break; + + size_t to_copy = bytes_read; + char* src = reinterpret_cast(threefs_iov.base); + while (to_copy > 0 && iov_idx < iovcnt) { + if (iov_off >= iov[iov_idx].iov_len) { + ++iov_idx; + iov_off = 0; + continue; + } + size_t n = std::min(to_copy, iov[iov_idx].iov_len - iov_off); + memcpy(static_cast(iov[iov_idx].iov_base) + iov_off, src, n); + src += n; + to_copy -= n; + total_read += n; + remaining -= n; + current_offset += n; + iov_off += n; + } + if (bytes_read < chunk) break; + } + + if (total_read != total_length) { + return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + } + return total_read; +} + } // namespace mooncake diff --git a/mooncake-store/src/storage/distributed/posix_fs_adapter.cpp b/mooncake-store/src/storage/distributed/posix_fs_adapter.cpp new file mode 100644 index 0000000000..bd20e0b467 --- /dev/null +++ b/mooncake-store/src/storage/distributed/posix_fs_adapter.cpp @@ -0,0 +1,236 @@ +#include "storage/distributed/posix_fs_adapter.h" + +#include +#include +#include +#include +#include + +#include +#include +#include + +namespace mooncake { + +namespace { + +tl::expected ValidateIov(const iovec* iov, int iovcnt) { + if (iovcnt < 0 || (iovcnt > 0 && iov == nullptr)) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + for (int i = 0; i < iovcnt; ++i) { + if (!iov[i].iov_base && iov[i].iov_len > 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + } + return {}; +} + +ErrorCode ReadOpenError() { + return errno == ENOENT ? ErrorCode::FILE_NOT_FOUND + : ErrorCode::FILE_OPEN_FAIL; +} + +} // namespace + +tl::expected PosixFsAdapter::Init( + const std::string& mount_path) { + mount_path_ = mount_path; + std::error_code ec; + std::filesystem::create_directories(mount_path_, ec); + if (ec) return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + return {}; +} + +tl::expected PosixFsAdapter::Shutdown() { return {}; } + +tl::expected PosixFsAdapter::WriteFile( + const std::string& path, std::span data) { + int fd = + ::open(path.c_str(), O_WRONLY | O_CREAT | O_TRUNC | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + + size_t total_written = 0; + while (total_written < data.size()) { + ssize_t ret = ::write(fd, data.data() + total_written, + data.size() - total_written); + if (ret < 0) { + int saved_errno = errno; + ::close(fd); + errno = saved_errno; + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + if (ret == 0) break; + total_written += static_cast(ret); + } + + if (::close(fd) != 0) { + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + if (total_written != data.size()) { + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + return total_written; +} + +tl::expected PosixFsAdapter::ReadFile( + const std::string& path, void* buf, size_t len) { + if (!buf && len > 0) return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + + int fd = ::open(path.c_str(), O_RDONLY | O_CLOEXEC); + if (fd < 0) return tl::make_unexpected(ReadOpenError()); + + size_t total_read = 0; + char* dest = static_cast(buf); + while (total_read < len) { + ssize_t ret = ::read(fd, dest + total_read, len - total_read); + if (ret < 0) { + int saved_errno = errno; + ::close(fd); + errno = saved_errno; + return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + } + if (ret == 0) break; + total_read += static_cast(ret); + } + + if (::close(fd) != 0) { + return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + } + return total_read; +} + +tl::expected PosixFsAdapter::VectorWriteFile( + const std::string& path, const iovec* iov, int iovcnt, off_t offset) { + auto valid = ValidateIov(iov, iovcnt); + if (!valid) return tl::make_unexpected(valid.error()); + if (offset < 0) return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + + int fd = + ::open(path.c_str(), O_WRONLY | O_CREAT | O_TRUNC | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + auto result = WriteAt(fd, iov, iovcnt, offset); + int saved_errno = errno; + ::close(fd); + errno = saved_errno; + return result; +} + +tl::expected PosixFsAdapter::VectorReadFile( + const std::string& path, const iovec* iov, int iovcnt, off_t offset) { + auto valid = ValidateIov(iov, iovcnt); + if (!valid) return tl::make_unexpected(valid.error()); + if (offset < 0) return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + + int fd = ::open(path.c_str(), O_RDONLY | O_CLOEXEC); + if (fd < 0) return tl::make_unexpected(ReadOpenError()); + auto result = ReadAt(fd, const_cast(iov), iovcnt, offset); + int saved_errno = errno; + ::close(fd); + errno = saved_errno; + return result; +} + +tl::expected PosixFsAdapter::DeleteFile( + const std::string& path) { + if (::unlink(path.c_str()) != 0) { + if (errno == ENOENT) + return tl::make_unexpected(ErrorCode::FILE_NOT_FOUND); + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + return {}; +} + +tl::expected PosixFsAdapter::FileExists( + const std::string& path) { + if (::access(path.c_str(), F_OK) == 0) return true; + if (errno == ENOENT) return false; + return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); +} + +tl::expected, ErrorCode> PosixFsAdapter::ListFiles( + const std::string& dir) { + DIR* d = ::opendir(dir.c_str()); + if (!d) { + if (errno == ENOENT) + return tl::make_unexpected(ErrorCode::FILE_NOT_FOUND); + return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + } + + std::vector result; + while (auto* entry = ::readdir(d)) { + std::string name = entry->d_name; + if (name == "." || name == "..") continue; + if (entry->d_type == DT_DIR) continue; + if (entry->d_type == DT_UNKNOWN) { + struct stat st; + if (::stat((dir + "/" + name).c_str(), &st) == 0 && + S_ISDIR(st.st_mode)) { + continue; + } + } + result.push_back(std::move(name)); + } + ::closedir(d); + return result; +} + +tl::expected PosixFsAdapter::OpenFile(const std::string& path) { + int fd = ::open(path.c_str(), O_RDWR | O_CREAT | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + return fd; +} + +tl::expected PosixFsAdapter::CloseFile(int fd) { + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_INVALID_HANDLE); + if (::close(fd) != 0) { + return tl::make_unexpected(ErrorCode::FILE_INVALID_HANDLE); + } + return {}; +} + +tl::expected PosixFsAdapter::PreallocateFile( + const std::string& path, uint64_t size) { + int fd = ::open(path.c_str(), O_RDWR | O_CREAT | O_CLOEXEC, 0644); + if (fd < 0) return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + + int rc = ::ftruncate(fd, static_cast(size)); + int saved_errno = errno; + ::close(fd); + errno = saved_errno; + if (rc != 0) return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + return {}; +} + +tl::expected PosixFsAdapter::WriteAt(int fd, + const iovec* iov, + int iovcnt, + int64_t offset) { + if (fd < 0 || offset < 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + auto valid = ValidateIov(iov, iovcnt); + if (!valid) return tl::make_unexpected(valid.error()); + if (iovcnt == 0) return size_t{0}; + + ssize_t ret = ::pwritev(fd, iov, iovcnt, static_cast(offset)); + if (ret < 0) return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + return static_cast(ret); +} + +tl::expected PosixFsAdapter::ReadAt(int fd, iovec* iov, + int iovcnt, + int64_t offset) { + if (fd < 0 || offset < 0) { + return tl::make_unexpected(ErrorCode::INVALID_PARAMS); + } + auto valid = ValidateIov(iov, iovcnt); + if (!valid) return tl::make_unexpected(valid.error()); + if (iovcnt == 0) return size_t{0}; + + ssize_t ret = ::preadv(fd, iov, iovcnt, static_cast(offset)); + if (ret < 0) return tl::make_unexpected(ErrorCode::FILE_READ_FAIL); + return static_cast(ret); +} + +} // namespace mooncake diff --git a/mooncake-store/src/storage_backend.cpp b/mooncake-store/src/storage_backend.cpp index 5a5e1588db..75dbd198d9 100644 --- a/mooncake-store/src/storage_backend.cpp +++ b/mooncake-store/src/storage_backend.cpp @@ -51,6 +51,10 @@ struct FdGuard { } // namespace #include "storage/distributed/distributed_storage_backend.h" +#include "storage/distributed/posix_fs_adapter.h" +#ifdef USE_3FS +#include "storage/distributed/hf3fs_adapter.h" +#endif namespace mooncake { @@ -5551,7 +5555,9 @@ CreateStorageBackend(const FileStorageConfig& config) { "Invalid DistributedStorage configuration"); } std::unique_ptr adapter; - if (distributed_config.fs_adapter_type == "hf3fs") { + if (distributed_config.fs_adapter_type == "posix") { + adapter = std::make_unique(); + } else if (distributed_config.fs_adapter_type == "hf3fs") { #ifdef USE_3FS adapter = std::make_unique(); #else diff --git a/mooncake-store/src/types.cpp b/mooncake-store/src/types.cpp index 7535abccbf..1956c024c4 100644 --- a/mooncake-store/src/types.cpp +++ b/mooncake-store/src/types.cpp @@ -50,6 +50,7 @@ const std::string& toString(ErrorCode errorCode) noexcept { {ErrorCode::UNAVAILABLE_IN_CURRENT_STATUS, "UNAVAILABLE_IN_CURRENT_STATUS"}, {ErrorCode::UNAVAILABLE_IN_CURRENT_MODE, "UNAVAILABLE_IN_CURRENT_MODE"}, + {ErrorCode::NOT_SUPPORTED, "NOT_SUPPORTED"}, {ErrorCode::FILE_NOT_FOUND, "FILE_NOT_FOUND"}, {ErrorCode::FILE_OPEN_FAIL, "FILE_OPEN_FAIL"}, {ErrorCode::FILE_READ_FAIL, "FILE_READ_FAIL"}, diff --git a/mooncake-store/tests/CMakeLists.txt b/mooncake-store/tests/CMakeLists.txt index 7e951d0d37..f7b5559e72 100644 --- a/mooncake-store/tests/CMakeLists.txt +++ b/mooncake-store/tests/CMakeLists.txt @@ -145,6 +145,16 @@ add_store_test(host_port_fix_test host_port_fix_test.cpp) add_store_test(client_metrics_test client_metrics_test.cpp) add_store_test(ssd_metrics_test ssd_metrics_test.cpp) add_store_test(serializer_test serializer_test.cpp) +add_store_test(dfs_posix_test dfs_posix_test.cpp) +add_store_test(dfs_sync_client_test dfs_sync_client_test.cpp) +add_test(NAME dfs_batch_read_checksum_test + COMMAND dfs_sync_client_test + --gtest_filter=DfsSyncClientTest.BatchGetVerifiesDfsChecksum) +set_tests_properties(dfs_batch_read_checksum_test + PROPERTIES ENVIRONMENT "MOONCAKE_STORE_CHECKSUM=1") +if(USE_3FS) + add_store_test(dfs_hf3fs_test dfs_hf3fs_test.cpp) +endif() add_store_test( embedded_snapshot_catalog_store_test ha/snapshot/catalog/backends/embedded/embedded_snapshot_catalog_store_test.cpp diff --git a/mooncake-store/tests/dfs_hf3fs_test.cpp b/mooncake-store/tests/dfs_hf3fs_test.cpp new file mode 100644 index 0000000000..2fb780d7e8 --- /dev/null +++ b/mooncake-store/tests/dfs_hf3fs_test.cpp @@ -0,0 +1,178 @@ +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "hf3fs/hf3fs.h" +#include "storage/distributed/dfs_global_allocator.h" +#include "storage/distributed/distributed_storage_backend.h" +#include "storage/distributed/hf3fs_adapter.h" +#include "storage_backend.h" + +namespace mooncake::test { +namespace { + +constexpr const char* kDefaultRoot = "/mnt/3fs/mooncake_test"; + +bool IsHf3fsPath(const std::string& path) { + char mount_point[PATH_MAX] = {}; + int ret = hf3fs_extract_mount_point(mount_point, sizeof(mount_point), + path.c_str()); + return ret > 0 && ret <= static_cast(sizeof(mount_point)); +} + +std::string DfsRootFromEnv() { + const char* root = std::getenv("MOONCAKE_DFS_ROOT_DIR"); + return root && root[0] != '\0' ? std::string(root) + : std::string(kDefaultRoot); +} + +std::optional ExistingAncestor(std::filesystem::path path) { + while (!path.empty()) { + if (std::filesystem::exists(path)) return path.string(); + path = path.parent_path(); + } + return std::nullopt; +} + +class Hf3fsTestDir { + public: + Hf3fsTestDir() { + static std::atomic counter{0}; + root_ = DfsRootFromEnv(); + path_ = std::filesystem::path(root_) / + ("dfs_hf3fs_test_" + std::to_string(::getpid()) + "_" + + std::to_string(++counter)); + path_str_ = path_.string(); + } + + ~Hf3fsTestDir() { + std::error_code ec; + std::filesystem::remove_all(path_, ec); + } + + void CreateOrSkip() { + auto ancestor = ExistingAncestor(path_); + if (!ancestor.has_value()) { + GTEST_SKIP() << "no existing parent for test root: " << root_; + } + if (!IsHf3fsPath(*ancestor)) { + GTEST_SKIP() << "test root is not under an hf3fs mount: " << root_; + } + + std::filesystem::create_directories(path_); + if (!IsHf3fsPath(path_str_)) { + GTEST_SKIP() << "test directory is not on an hf3fs mount: " + << path_str_; + } + } + + const std::string& path() const { return path_str_; } + + std::string file(const std::string& name) const { + return (path_ / name).string(); + } + + private: + std::string root_; + std::filesystem::path path_; + std::string path_str_; +}; + +class Hf3fsAdapterTest : public ::testing::Test { + protected: + void SetUp() override { + test_dir_ = std::make_unique(); + test_dir_->CreateOrSkip(); + } + + void TearDown() override { test_dir_.reset(); } + + std::unique_ptr test_dir_; +}; + +} // namespace + +TEST_F(Hf3fsAdapterTest, WriteAtReadAtThroughUsrbio) { + Hf3fsAdapter adapter; + ASSERT_TRUE(adapter.Init(test_dir_->path()).has_value()); + + const std::string path = test_dir_->file("adapter_smoke.data"); + ASSERT_TRUE(adapter.PreallocateFile(path, 4096).has_value()); + + auto fd = adapter.OpenFile(path); + ASSERT_TRUE(fd.has_value()); + + std::array write_buf; + std::array read_buf{}; + write_buf.fill('Q'); + + iovec wiov{write_buf.data(), write_buf.size()}; + auto written = adapter.WriteAt(*fd, &wiov, 1, 100); + ASSERT_TRUE(written.has_value()); + EXPECT_EQ(*written, write_buf.size()); + + iovec riov{read_buf.data(), read_buf.size()}; + auto read = adapter.ReadAt(*fd, &riov, 1, 100); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, read_buf.size()); + EXPECT_EQ(std::memcmp(write_buf.data(), read_buf.data(), write_buf.size()), + 0); + + EXPECT_TRUE(adapter.CloseFile(*fd).has_value()); + EXPECT_TRUE(adapter.Shutdown().has_value()); +} + +TEST_F(Hf3fsAdapterTest, DistributedBackendBatchWriteAndRead) { + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = test_dir_->path(); + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = test_dir_->path(); + distributed_config.fs_adapter_type = "hf3fs"; + distributed_config.shard_count = 2; + distributed_config.shard_capacity = 1024 * 1024; + distributed_config.alignment = 4096; + distributed_config.single_tenant = true; + + DistributedStorageBackend backend(file_config, distributed_config, + std::make_unique()); + ASSERT_TRUE(backend.Init().has_value()); + + alignas(4096) std::array write_buf; + alignas(4096) std::array read_buf{}; + write_buf.fill('B'); + + const std::string key = "hf3fs_backend_key"; + const std::string shard_path = test_dir_->file( + "dfs_shard_" + + DfsGlobalAllocator::FormatShardIdx(0, distributed_config.shard_count) + + ".data"); + DistributedFSDescriptor descriptor{shard_path, 0, write_buf.size(), + write_buf.size(), 0}; + auto write = backend.BatchWrite( + {{key, descriptor, {{write_buf.data(), write_buf.size()}}}}); + ASSERT_EQ(write.size(), 1); + ASSERT_TRUE(write[0].has_value()); + + auto read = backend.BatchRead( + {{key, descriptor, {{read_buf.data(), read_buf.size()}}}}); + ASSERT_EQ(read.size(), 1); + ASSERT_TRUE(read[0].has_value()); + EXPECT_EQ(std::memcmp(write_buf.data(), read_buf.data(), write_buf.size()), + 0); +} + +} // namespace mooncake::test diff --git a/mooncake-store/tests/dfs_posix_test.cpp b/mooncake-store/tests/dfs_posix_test.cpp new file mode 100644 index 0000000000..1e5ec533d6 --- /dev/null +++ b/mooncake-store/tests/dfs_posix_test.cpp @@ -0,0 +1,1166 @@ +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "replica.h" +#include "storage/distributed/dfs_global_allocator.h" +#include "storage/distributed/distributed_storage_backend.h" +#include "storage/distributed/posix_fs_adapter.h" +#include "storage_backend.h" + +namespace mooncake::test { + +class TempDir { + public: + explicit TempDir(const std::string& prefix) { + static std::atomic counter{0}; + path_ = std::filesystem::temp_directory_path() / + (prefix + "_" + std::to_string(::getpid()) + "_" + + std::to_string(++counter)); + path_str_ = path_.string(); + std::filesystem::create_directories(path_); + } + + ~TempDir() { + std::error_code ec; + std::filesystem::remove_all(path_, ec); + } + + const std::string& path() const { return path_str_; } + + std::string file(const std::string& name) const { + return (path_ / name).string(); + } + + private: + std::filesystem::path path_; + std::string path_str_; +}; + +class EnvGuard { + public: + EnvGuard() { + Save("MOONCAKE_DFS_FS_ADAPTER"); + Save("MOONCAKE_DISTRIBUTED_FS_TYPE"); + Save("MOONCAKE_DFS_EVICTION_ENABLED"); + Save("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK"); + Save("MOONCAKE_DFS_EVICTION_LOW_WATERMARK"); + Save("MOONCAKE_DFS_DEFERRED_FREE_SECONDS"); + Save("MOONCAKE_DFS_EVICTION_CHECK_INTERVAL"); + Save("MOONCAKE_DFS_ROOT_DIR"); + Save("MOONCAKE_DFS_SHARD_COUNT"); + Save("MOONCAKE_DFS_SHARD_CAPACITY"); + Save("MOONCAKE_DFS_ALIGNMENT"); + Save("MOONCAKE_DFS_SINGLE_TENANT"); + } + + ~EnvGuard() { + for (const auto& [key, value] : saved_) { + if (value.has_value()) { + ::setenv(key.c_str(), value->c_str(), 1); + } else { + ::unsetenv(key.c_str()); + } + } + } + + void Set(const char* key, const char* value) { ::setenv(key, value, 1); } + + private: + void Save(const std::string& key) { + const char* value = ::getenv(key.c_str()); + if (value) { + saved_.push_back({key, std::string(value)}); + } else { + saved_.push_back({key, std::nullopt}); + } + } + + std::vector>> saved_; +}; + +class AlignedBuffer { + public: + explicit AlignedBuffer(size_t size, size_t alignment = 4096) : size_(size) { + void* ptr = nullptr; + if (::posix_memalign(&ptr, alignment, size) != 0) ptr = nullptr; + ptr_ = static_cast(ptr); + } + + ~AlignedBuffer() { std::free(ptr_); } + + AlignedBuffer(const AlignedBuffer&) = delete; + AlignedBuffer& operator=(const AlignedBuffer&) = delete; + + char* data() { return ptr_; } + const char* data() const { return ptr_; } + size_t size() const { return size_; } + + void Fill(char value) { std::memset(ptr_, value, size_); } + + private: + char* ptr_ = nullptr; + size_t size_ = 0; +}; + +void ConfigurePosixDfs(EnvGuard& env) { + env.Set("MOONCAKE_DFS_FS_ADAPTER", "posix"); + env.Set("MOONCAKE_DFS_SINGLE_TENANT", "1"); + env.Set("MOONCAKE_DFS_EVICTION_ENABLED", "0"); + env.Set("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "0"); +} + +DistributedStorageConfig MakeAllocatorConfig(const std::string& mount_path, + int shard_count, + uint64_t shard_capacity, + uint64_t alignment) { + auto config = DistributedStorageConfig::FromEnvironment(); + config.fsdir = mount_path; + config.shard_count = shard_count; + config.shard_capacity = shard_capacity; + config.alignment = alignment; + return config; +} + +std::vector +PrepareAndCommitPreparedEviction(DfsGlobalAllocator& allocator) { + auto pending = allocator.PrepareEviction(); + auto candidates = pending.Candidates(); + allocator.CommitPreparedEviction(std::move(pending)); + return candidates; +} + +class FsAdapterFdTest : public ::testing::Test { + protected: + void SetUp() override { + tmp_ = std::make_unique("dfs_fd_test"); + adapter_ = std::make_unique(); + ASSERT_TRUE(adapter_->Init(tmp_->path()).has_value()); + } + + void TearDown() override { + adapter_.reset(); + tmp_.reset(); + } + + std::unique_ptr tmp_; + std::unique_ptr adapter_; +}; + +TEST_F(FsAdapterFdTest, OpenClose) { + auto pre = adapter_->PreallocateFile(tmp_->file("shard0.data"), 4096); + ASSERT_TRUE(pre.has_value()); + + auto fd = adapter_->OpenFile(tmp_->file("shard0.data")); + ASSERT_TRUE(fd.has_value()); + EXPECT_GE(*fd, 0); + + auto close = adapter_->CloseFile(*fd); + EXPECT_TRUE(close.has_value()); +} + +TEST_F(FsAdapterFdTest, WriteAtReadAt) { + ASSERT_TRUE( + adapter_->PreallocateFile(tmp_->file("shard0.data"), 4096).has_value()); + auto fd = adapter_->OpenFile(tmp_->file("shard0.data")); + ASSERT_TRUE(fd.has_value()); + + char write_buf[128]; + std::memset(write_buf, 'A', sizeof(write_buf)); + iovec wiov{write_buf, sizeof(write_buf)}; + auto written = adapter_->WriteAt(*fd, &wiov, 1, 100); + ASSERT_TRUE(written.has_value()); + EXPECT_EQ(*written, sizeof(write_buf)); + + char read_buf[128] = {}; + iovec riov{read_buf, sizeof(read_buf)}; + auto read = adapter_->ReadAt(*fd, &riov, 1, 100); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, sizeof(read_buf)); + EXPECT_EQ(std::memcmp(write_buf, read_buf, sizeof(write_buf)), 0); + + adapter_->CloseFile(*fd); +} + +TEST_F(FsAdapterFdTest, MultiIovWriteRead) { + constexpr size_t total = 8192; + ASSERT_TRUE(adapter_->PreallocateFile(tmp_->file("shard_multi.data"), total) + .has_value()); + auto fd = adapter_->OpenFile(tmp_->file("shard_multi.data")); + ASSERT_TRUE(fd.has_value()); + + char w0[2048], w1[3072], w2[3072]; + std::memset(w0, 'A', sizeof(w0)); + std::memset(w1, 'B', sizeof(w1)); + std::memset(w2, 'C', sizeof(w2)); + iovec wiovs[3] = { + {w0, sizeof(w0)}, + {w1, sizeof(w1)}, + {w2, sizeof(w2)}, + }; + auto written = adapter_->WriteAt(*fd, wiovs, 3, 0); + ASSERT_TRUE(written.has_value()); + EXPECT_EQ(*written, total); + + char r0[2048] = {}, r1[3072] = {}, r2[3072] = {}; + iovec riovs[3] = { + {r0, sizeof(r0)}, + {r1, sizeof(r1)}, + {r2, sizeof(r2)}, + }; + auto read = adapter_->ReadAt(*fd, riovs, 3, 0); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, total); + EXPECT_EQ(std::memcmp(w0, r0, sizeof(w0)), 0); + EXPECT_EQ(std::memcmp(w1, r1, sizeof(w1)), 0); + EXPECT_EQ(std::memcmp(w2, r2, sizeof(w2)), 0); + + adapter_->CloseFile(*fd); +} + +TEST_F(FsAdapterFdTest, MultiIovPartialReadAndUnalignedAccess) { + constexpr size_t total = 8192; + ASSERT_TRUE( + adapter_->PreallocateFile(tmp_->file("shard_partial.data"), total) + .has_value()); + auto fd = adapter_->OpenFile(tmp_->file("shard_partial.data")); + ASSERT_TRUE(fd.has_value()); + + std::string write_data(total, 'X'); + iovec wiov{write_data.data(), write_data.size()}; + ASSERT_TRUE(adapter_->WriteAt(*fd, &wiov, 1, 0).has_value()); + + char read_buf[3072] = {}; + iovec riov{read_buf, sizeof(read_buf)}; + auto read = adapter_->ReadAt(*fd, &riov, 1, 2048); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, sizeof(read_buf)); + EXPECT_EQ(std::memcmp(read_buf, write_data.data() + 2048, sizeof(read_buf)), + 0); + + char wbuf[63]; + std::memset(wbuf, 'Y', sizeof(wbuf)); + iovec unaligned_wiov{wbuf, sizeof(wbuf)}; + auto written = adapter_->WriteAt(*fd, &unaligned_wiov, 1, 101); + ASSERT_TRUE(written.has_value()); + EXPECT_EQ(*written, sizeof(wbuf)); + + char rbuf[63] = {}; + iovec unaligned_riov{rbuf, sizeof(rbuf)}; + read = adapter_->ReadAt(*fd, &unaligned_riov, 1, 101); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, sizeof(rbuf)); + EXPECT_EQ(std::memcmp(wbuf, rbuf, sizeof(wbuf)), 0); + + char beyond[64] = {}; + iovec beyond_iov{beyond, sizeof(beyond)}; + auto beyond_read = adapter_->ReadAt(*fd, &beyond_iov, 1, 1ULL << 30); + ASSERT_TRUE(beyond_read.has_value()); + EXPECT_EQ(*beyond_read, 0); + + adapter_->CloseFile(*fd); +} + +TEST_F(FsAdapterFdTest, PreallocateLargeSparseFile) { + constexpr uint64_t size = 4ULL * 1024 * 1024 * 1024; + ASSERT_TRUE(adapter_->PreallocateFile(tmp_->file("shard_large.data"), size) + .has_value()); + + struct stat st; + ASSERT_EQ(::stat(tmp_->file("shard_large.data").c_str(), &st), 0); + EXPECT_EQ(static_cast(st.st_size), size); + EXPECT_LT(static_cast(st.st_blocks) * 512, size / 1000); + + auto fd = adapter_->OpenFile(tmp_->file("shard_large.data")); + ASSERT_TRUE(fd.has_value()); + char wbuf[4096]; + std::memset(wbuf, 'L', sizeof(wbuf)); + iovec wiov{wbuf, sizeof(wbuf)}; + auto written = adapter_->WriteAt(*fd, &wiov, 1, 3ULL * 1024 * 1024 * 1024); + ASSERT_TRUE(written.has_value()); + EXPECT_EQ(*written, sizeof(wbuf)); + + char rbuf[4096] = {}; + iovec riov{rbuf, sizeof(rbuf)}; + auto read = adapter_->ReadAt(*fd, &riov, 1, 3ULL * 1024 * 1024 * 1024); + ASSERT_TRUE(read.has_value()); + EXPECT_EQ(*read, sizeof(rbuf)); + EXPECT_EQ(std::memcmp(wbuf, rbuf, sizeof(wbuf)), 0); + + adapter_->CloseFile(*fd); +} + +TEST(DfsGlobalAllocatorTest, AllocateFreeAndFormatShardIdx) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_alloc"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 4, 1024 * 1024, 4096))); + + auto desc = alloc.Allocate("key1", 100); + ASSERT_TRUE(desc.has_value()); + EXPECT_EQ(desc->aligned_size, 4096); + EXPECT_GE(desc->shard_idx, 0); + EXPECT_LT(desc->shard_idx, 4); + EXPECT_EQ(desc->offset % 4096, 0); + + alloc.Free(desc->offset, desc->aligned_size, desc->shard_idx, "key1"); + auto desc2 = alloc.Allocate("key2", 100); + EXPECT_TRUE(desc2.has_value()); + + EXPECT_EQ(DfsGlobalAllocator::FormatShardIdx(0, 64), "00"); + EXPECT_EQ(DfsGlobalAllocator::FormatShardIdx(9, 64), "09"); + EXPECT_EQ(DfsGlobalAllocator::FormatShardIdx(10, 64), "10"); + EXPECT_EQ(DfsGlobalAllocator::FormatShardIdx(63, 64), "63"); + EXPECT_EQ(DfsGlobalAllocator::FormatShardIdx(100, 1000), "100"); +} + +TEST(DistributedStorageConfigTest, ReadsValidatesAndFormatsEnvironment) { + EnvGuard env; + TempDir tmp("dfs_config"); + env.Set("MOONCAKE_DFS_ROOT_DIR", tmp.path().c_str()); + env.Set("MOONCAKE_DFS_FS_ADAPTER", "posix"); + env.Set("MOONCAKE_DFS_SHARD_COUNT", "8"); + env.Set("MOONCAKE_DFS_SHARD_CAPACITY", "1048576"); + env.Set("MOONCAKE_DFS_ALIGNMENT", "4096"); + env.Set("MOONCAKE_DFS_SINGLE_TENANT", "1"); + env.Set("MOONCAKE_DFS_EVICTION_ENABLED", "1"); + env.Set("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", "0.85"); + env.Set("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", "0.65"); + env.Set("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "12"); + env.Set("MOONCAKE_DFS_EVICTION_CHECK_INTERVAL", "3"); + + const auto config = DistributedStorageConfig::FromEnvironment(); + EXPECT_EQ(config.fsdir, tmp.path()); + EXPECT_EQ(config.fs_adapter_type, "posix"); + EXPECT_EQ(config.shard_count, 8); + EXPECT_EQ(config.shard_capacity, 1048576); + EXPECT_EQ(config.alignment, 4096); + EXPECT_TRUE(config.eviction_enabled); + EXPECT_DOUBLE_EQ(config.eviction_high_watermark, 0.85); + EXPECT_DOUBLE_EQ(config.eviction_low_watermark, 0.65); + EXPECT_EQ(config.deferred_free_duration, std::chrono::seconds(12)); + EXPECT_EQ(config.eviction_check_interval, std::chrono::seconds(3)); + EXPECT_TRUE(config.Validate()); + EXPECT_TRUE(config.ValidateForAllocator()); + + const std::string formatted = config.FormatStr(); + EXPECT_NE(formatted.find("fs_adapter_type=posix"), std::string::npos); + EXPECT_NE(formatted.find("shard_count=8"), std::string::npos); + EXPECT_NE(formatted.find("eviction_high_watermark=0.85"), + std::string::npos); + + auto invalid_eviction = config; + invalid_eviction.eviction_low_watermark = 0.9; + EXPECT_TRUE(invalid_eviction.Validate()); + EXPECT_FALSE(invalid_eviction.ValidateForAllocator()); +} + +TEST(DfsGlobalAllocatorTest, InitReturnsSpecificErrors) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_init_error"); + + DistributedStorageConfig invalid_config = + MakeAllocatorConfig(tmp.path(), 1, 1024 * 1024, 3); + DfsGlobalAllocator invalid_allocator; + auto invalid_result = invalid_allocator.Init(invalid_config); + ASSERT_FALSE(invalid_result); + EXPECT_EQ(invalid_result.error(), ErrorCode::INVALID_PARAMS); + + const std::string file_path = tmp.file("not_a_directory"); + const int fd = ::open(file_path.c_str(), O_CREAT | O_WRONLY, 0600); + ASSERT_GE(fd, 0); + ASSERT_EQ(::close(fd), 0); + + DistributedStorageConfig file_error_config = + MakeAllocatorConfig(file_path, 1, 1024 * 1024, 4096); + DfsGlobalAllocator file_error_allocator; + auto file_error_result = file_error_allocator.Init(file_error_config); + ASSERT_FALSE(file_error_result); + EXPECT_EQ(file_error_result.error(), ErrorCode::FILE_WRITE_FAIL); +} + +TEST(DfsGlobalAllocatorTest, AllocateReservesAlignmentPadding) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_alloc_padding"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE(alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 8 * 1024, 4096))); + + auto desc = alloc.Allocate("key1", 100); + ASSERT_TRUE(desc.has_value()); + EXPECT_EQ(desc->aligned_size, 4096); + EXPECT_EQ(desc->offset % 4096, 0); + + auto exhausted = alloc.Allocate("key2", 100); + EXPECT_FALSE(exhausted.has_value()); + EXPECT_EQ(exhausted.error(), ErrorCode::NO_AVAILABLE_HANDLE); + + alloc.Free(desc->offset, desc->aligned_size, desc->shard_idx, "key1"); + auto after_free = alloc.Allocate("key3", 100); + EXPECT_TRUE(after_free.has_value()); +} + +TEST(DfsGlobalAllocatorTest, ExhaustionAndEviction) { + EnvGuard env; + ConfigurePosixDfs(env); + env.Set("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", "0.5"); + env.Set("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", "0.25"); + TempDir tmp("dfs_exhaust"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 32 * 1024, 4096))); + + std::vector descs; + for (int i = 0; i < 4; ++i) { + auto desc = alloc.Allocate("k" + std::to_string(i), 100); + ASSERT_TRUE(desc.has_value()) << "allocation " << i; + alloc.UpdateAccess("k" + std::to_string(i), desc->shard_idx, + desc->offset); + descs.push_back(*desc); + } + + auto exhausted = alloc.Allocate("k_exhausted", 100); + EXPECT_FALSE(exhausted.has_value()); + EXPECT_EQ(exhausted.error(), ErrorCode::NO_AVAILABLE_HANDLE); + + auto pending = alloc.PrepareEviction(); + auto evicted = pending.Candidates(); + ASSERT_FALSE(evicted.empty()); + EXPECT_EQ(evicted.front().key, "k0"); + + // Prepare only reserves candidates; their extents are not reusable until + // the master accepts and commits the transaction. + auto before_commit = alloc.Allocate("k_before_commit", 100); + EXPECT_FALSE(before_commit.has_value()); + + alloc.CommitPreparedEviction(std::move(pending)); + + auto after_evict = alloc.Allocate("k_after_evict", 100); + EXPECT_TRUE(after_evict.has_value()); +} + +TEST(DfsGlobalAllocatorTest, RestorePreparedEvictionPreservesCandidateOrder) { + EnvGuard env; + ConfigurePosixDfs(env); + env.Set("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", "0.9"); + env.Set("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", "0.7"); + TempDir tmp("dfs_abort_eviction"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 32 * 1024, 4096))); + + for (int i = 0; i < 4; ++i) { + const std::string key = "k" + std::to_string(i); + auto desc = alloc.Allocate(key, 100); + ASSERT_TRUE(desc.has_value()) << "allocation " << i; + alloc.UpdateAccess(key, desc->shard_idx, desc->offset); + } + + auto pending = alloc.PrepareEviction(); + ASSERT_EQ(pending.Candidates().size(), 2); + EXPECT_EQ(pending.Candidates()[0].key, "k0"); + EXPECT_EQ(pending.Candidates()[1].key, "k1"); + alloc.RestorePreparedEviction(std::move(pending)); + + auto retry = alloc.PrepareEviction(); + ASSERT_EQ(retry.Candidates().size(), 2); + EXPECT_EQ(retry.Candidates()[0].key, "k0"); + EXPECT_EQ(retry.Candidates()[1].key, "k1"); + alloc.RestorePreparedEviction(std::move(retry)); +} + +TEST(DfsGlobalAllocatorTest, PartialResolutionContinuesToLowWatermark) { + EnvGuard env; + ConfigurePosixDfs(env); + env.Set("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", "0.9"); + env.Set("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", "0.7"); + TempDir tmp("dfs_partial_eviction"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 32 * 1024, 4096))); + + std::vector descs; + for (int i = 0; i < 4; ++i) { + const std::string key = "k" + std::to_string(i); + auto desc = alloc.Allocate(key, 100); + ASSERT_TRUE(desc.has_value()) << "allocation " << i; + alloc.UpdateAccess(key, desc->shard_idx, desc->offset); + descs.push_back(*desc); + } + + auto pending = alloc.PrepareEviction(); + ASSERT_EQ(pending.Candidates().size(), 2); + EXPECT_EQ(pending.Candidates()[0].key, "k0"); + EXPECT_EQ(pending.Candidates()[1].key, "k1"); + + alloc.ResolvePreparedEviction(std::move(pending), {true, false}); + alloc.UpdateAccess("k1", descs[1].shard_idx, descs[1].offset); + + // Accepting only k0 drops usage below the high watermark but not the low + // watermark. The same active eviction cycle must therefore continue and + // reach k2 behind the restored, protected k1. + auto continuation = alloc.PrepareEviction(); + ASSERT_EQ(continuation.Candidates().size(), 1); + EXPECT_EQ(continuation.Candidates().front().key, "k2"); + alloc.CommitPreparedEviction(std::move(continuation)); + + auto complete = alloc.PrepareEviction(); + EXPECT_TRUE(complete.Empty()); +} + +TEST(DfsGlobalAllocatorTest, EvictionCountsPendingFreeTowardWatermarks) { + EnvGuard env; + ConfigurePosixDfs(env); + env.Set("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", "0.9"); + env.Set("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", "0.7"); + env.Set("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "30"); + TempDir tmp("dfs_pending_watermark"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 32 * 1024, 4096))); + + for (int i = 0; i < 4; ++i) { + const std::string key = "k" + std::to_string(i); + auto desc = alloc.Allocate(key, 100); + ASSERT_TRUE(desc.has_value()) << "allocation " << i; + alloc.UpdateAccess(key, desc->shard_idx, desc->offset); + } + + auto evicted = PrepareAndCommitPreparedEviction(alloc); + ASSERT_EQ(evicted.size(), 2); + EXPECT_EQ(evicted[0].key, "k0"); + EXPECT_EQ(evicted[1].key, "k1"); + + auto repeated = PrepareAndCommitPreparedEviction(alloc); + EXPECT_TRUE(repeated.empty()); + + auto before_release = alloc.Allocate("k_before_release", 100); + EXPECT_FALSE(before_release.has_value()); + EXPECT_EQ(before_release.error(), ErrorCode::NO_AVAILABLE_HANDLE); +} + +TEST(DfsGlobalAllocatorTest, FreeRemovesLruEntryBeforeOffsetReuse) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_free_lru_reuse"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE(alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 8 * 1024, 4096))); + + auto desc_a = alloc.Allocate("A", 100); + ASSERT_TRUE(desc_a.has_value()); + alloc.UpdateAccess("A", desc_a->shard_idx, desc_a->offset); + + alloc.Free(desc_a->offset, desc_a->aligned_size, desc_a->shard_idx, "A"); + + auto desc_b = alloc.Allocate("B", 100); + ASSERT_TRUE(desc_b.has_value()); + ASSERT_EQ(desc_b->offset, desc_a->offset); + alloc.UpdateAccess("B", desc_b->shard_idx, desc_b->offset); + + auto evicted = PrepareAndCommitPreparedEviction(alloc); + ASSERT_EQ(evicted.size(), 1); + EXPECT_EQ(evicted.front().key, "B"); +} + +TEST(DfsGlobalAllocatorTest, StaleFreeDoesNotReleaseReusedOffset) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_stale_free"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE(alloc.Init(MakeAllocatorConfig(tmp.path(), 1, 8 * 1024, 4096))); + + auto desc_a = alloc.Allocate("A", 100); + ASSERT_TRUE(desc_a.has_value()); + alloc.UpdateAccess("A", desc_a->shard_idx, desc_a->offset); + + alloc.Free(desc_a->offset, desc_a->aligned_size, desc_a->shard_idx, "A"); + + auto desc_b = alloc.Allocate("B", 100); + ASSERT_TRUE(desc_b.has_value()); + ASSERT_EQ(desc_b->offset, desc_a->offset); + alloc.UpdateAccess("B", desc_b->shard_idx, desc_b->offset); + + alloc.Free(desc_a->offset, desc_a->aligned_size, desc_a->shard_idx, "A"); + + auto evicted = PrepareAndCommitPreparedEviction(alloc); + ASSERT_EQ(evicted.size(), 1); + EXPECT_EQ(evicted.front().key, "B"); +} + +TEST(DfsGlobalAllocatorTest, ConcurrentAllocate) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_concurrent"); + + DfsGlobalAllocator alloc; + ASSERT_TRUE( + alloc.Init(MakeAllocatorConfig(tmp.path(), 4, 128 * 1024, 4096))); + + constexpr int kThreadCount = 32; + std::vector threads; + std::atomic success_count{0}; + std::atomic fail_count{0}; + for (int i = 0; i < kThreadCount; ++i) { + threads.emplace_back([&alloc, &success_count, &fail_count, i]() { + std::string key = "key_" + std::to_string(i); + auto desc = alloc.Allocate(key, 100); + if (desc.has_value()) { + success_count++; + alloc.UpdateAccess(key, desc->shard_idx, desc->offset); + alloc.Free(desc->offset, desc->aligned_size, desc->shard_idx, + key); + } else { + fail_count++; + } + }); + } + for (auto& thread : threads) thread.join(); + + EXPECT_EQ(success_count.load(), kThreadCount); + EXPECT_EQ(fail_count.load(), 0); +} + +TEST(ReplicaDfsTest, HelpersAndDescriptor) { + DistributedFSDescriptor desc{"/mnt/3fs/shard0.data", 4096, 100, 4096, 0}; + Replica replica(desc, ReplicaStatus::PROCESSING); + + EXPECT_TRUE(replica.is_dfs_replica()); + EXPECT_FALSE(replica.is_memory_replica()); + EXPECT_FALSE(replica.is_disk_replica()); + EXPECT_FALSE(replica.is_nof_replica()); + EXPECT_EQ(replica.type(), ReplicaType::DFS); + EXPECT_EQ(replica.get_dfs_descriptor().offset, 4096); + + auto descriptor = replica.get_descriptor(); + EXPECT_TRUE(descriptor.is_dfs_replica()); + EXPECT_FALSE(descriptor.is_memory_replica()); + EXPECT_EQ(descriptor.status, ReplicaStatus::PROCESSING); + EXPECT_EQ(descriptor.get_dfs_descriptor().object_size, 100); + + EXPECT_TRUE(replica.is_processing()); + replica.mark_complete(); + EXPECT_TRUE(replica.is_completed()); + replica.mark_processing(); + EXPECT_TRUE(replica.is_processing()); + + EXPECT_EQ(replica.get_refcnt(), 0); + replica.inc_refcnt(); + EXPECT_TRUE(replica.is_busy()); + replica.dec_refcnt(); + EXPECT_FALSE(replica.is_busy()); + + ReplicateConfig config; + config.replica_num = 1; + config.nof_replica_num = 0; + config.dfs_replica_num = 1; + EXPECT_EQ(DetermineReplicaWriteMode(config), + ReplicaWriteMode::RELIABLE_MULTI_REPLICA); +} + +class DfsBackendTest : public ::testing::Test { + protected: + void SetUp() override { + tmp_ = std::make_unique("dfs_backend"); + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = tmp_->path(); + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = tmp_->path(); + distributed_config.fs_adapter_type = "posix"; + distributed_config.shard_count = 4; + distributed_config.shard_capacity = 64 * 1024 * 1024; + distributed_config.alignment = 4096; + + backend_ = std::make_unique( + file_config, distributed_config, + std::make_unique()); + ASSERT_TRUE(backend_->Init().has_value()); + } + + void TearDown() override { + backend_.reset(); + tmp_.reset(); + } + + std::string ShardPath(int shard_idx) const { + return tmp_->file("dfs_shard_" + + DfsGlobalAllocator::FormatShardIdx(shard_idx, 4) + + ".data"); + } + + std::unique_ptr tmp_; + std::unique_ptr backend_; +}; + +class ControlledPosixFsAdapter : public PosixFsAdapter { + public: + void FailWriteCall(int call) { fail_write_call_ = call; } + void ShortWriteCall(int call) { short_write_call_ = call; } + void FailReadCall(int call) { fail_read_call_ = call; } + void ShortReadCall(int call) { short_read_call_ = call; } + int WriteCallCount() const { return write_calls_.load(); } + int ReadCallCount() const { return read_calls_.load(); } + + tl::expected WriteAt(int fd, const iovec* iov, + int iovcnt, + int64_t offset) override { + const int call = ++write_calls_; + if (call == fail_write_call_) { + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + if (call == short_write_call_) { + size_t total_size = 0; + for (int i = 0; i < iovcnt; ++i) { + total_size += iov[i].iov_len; + } + return total_size == 0 ? 0 : total_size - 1; + } + return PosixFsAdapter::WriteAt(fd, iov, iovcnt, offset); + } + + tl::expected ReadAt(int fd, iovec* iov, int iovcnt, + int64_t offset) override { + const int call = ++read_calls_; + if (call == fail_read_call_) { + return tl::make_unexpected(ErrorCode::FILE_OPEN_FAIL); + } + if (call == short_read_call_) { + size_t total_size = 0; + for (int i = 0; i < iovcnt; ++i) { + total_size += iov[i].iov_len; + } + return total_size == 0 ? 0 : total_size - 1; + } + return PosixFsAdapter::ReadAt(fd, iov, iovcnt, offset); + } + + private: + std::atomic write_calls_{0}; + std::atomic read_calls_{0}; + int fail_write_call_ = -1; + int short_write_call_ = -1; + int fail_read_call_ = -1; + int short_read_call_ = -1; +}; + +TEST_F(DfsBackendTest, BatchWriteUsesExplicitDescriptors) { + AlignedBuffer write_buf(4096); + ASSERT_NE(write_buf.data(), nullptr); + write_buf.Fill('E'); + + std::vector requests{ + {"explicit", + {ShardPath(1), 0, 4096, 4096, 1}, + {{write_buf.data(), write_buf.size()}}}, + {"bad_path", + {tmp_->file("wrong.data"), 4096, 4096, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}, + {"bad_size", + {ShardPath(0), 4096, 2048, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}, + }; + auto results = backend_->BatchWrite(requests); + ASSERT_EQ(results.size(), 3); + EXPECT_TRUE(results[0].has_value()); + ASSERT_FALSE(results[1].has_value()); + EXPECT_EQ(results[1].error(), ErrorCode::INVALID_PARAMS); + ASSERT_FALSE(results[2].has_value()); + EXPECT_EQ(results[2].error(), ErrorCode::INVALID_PARAMS); + + AlignedBuffer read_buf(4096); + auto read_results = + backend_->BatchRead({{"explicit", + requests[0].descriptor, + {{read_buf.data(), read_buf.size()}}}}); + ASSERT_EQ(read_results.size(), 1); + ASSERT_TRUE(read_results[0].has_value()); + EXPECT_EQ(std::memcmp(write_buf.data(), read_buf.data(), write_buf.size()), + 0); +} + +TEST_F(DfsBackendTest, BatchWritePreservesPerKeyWriteErrors) { + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = tmp_->path(); + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = tmp_->path(); + distributed_config.fs_adapter_type = "posix"; + distributed_config.shard_count = 4; + distributed_config.shard_capacity = 64 * 1024 * 1024; + distributed_config.alignment = 4096; + + auto adapter = std::make_unique(); + auto* controlled_adapter = adapter.get(); + auto backend = std::make_unique( + file_config, distributed_config, std::move(adapter)); + ASSERT_TRUE(backend->Init().has_value()); + controlled_adapter->FailWriteCall(2); + controlled_adapter->ShortWriteCall(3); + + AlignedBuffer write_buf(4096); + ASSERT_NE(write_buf.data(), nullptr); + std::vector requests{ + {"ok", + {ShardPath(0), 0, 4096, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}, + {"failed", + {ShardPath(0), 4096, 4096, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}, + {"short", + {ShardPath(0), 8192, 4096, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}, + }; + auto results = backend->BatchWrite(requests); + ASSERT_EQ(results.size(), 3); + EXPECT_TRUE(results[0].has_value()); + ASSERT_FALSE(results[1].has_value()); + EXPECT_EQ(results[1].error(), ErrorCode::FILE_WRITE_FAIL); + ASSERT_FALSE(results[2].has_value()); + EXPECT_EQ(results[2].error(), ErrorCode::FILE_WRITE_FAIL); +} + +TEST_F(DfsBackendTest, BatchReadUsesExplicitDescriptorsAndMultipleSlices) { + AlignedBuffer first_value(4096), second_value(4096); + ASSERT_NE(first_value.data(), nullptr); + ASSERT_NE(second_value.data(), nullptr); + first_value.Fill('R'); + second_value.Fill('S'); + + const DistributedFSDescriptor first_desc{ShardPath(0), 0, 4096, 4096, 0}; + const DistributedFSDescriptor second_desc{ShardPath(0), 4096, 4096, 4096, + 0}; + auto write_results = backend_->BatchWrite( + {{"same_key", first_desc, {{first_value.data(), first_value.size()}}}, + {"other_key", + second_desc, + {{second_value.data(), second_value.size()}}}}); + ASSERT_EQ(write_results.size(), 2); + ASSERT_TRUE(write_results[0].has_value()); + ASSERT_TRUE(write_results[1].has_value()); + + AlignedBuffer output_0(1024), output_1(3072), unused(64); + ASSERT_NE(output_0.data(), nullptr); + ASSERT_NE(output_1.data(), nullptr); + ASSERT_NE(unused.data(), nullptr); + unused.Fill('U'); + std::vector requests{ + {"same_key", + first_desc, + {{output_0.data(), output_0.size()}, + {output_1.data(), output_1.size()}, + {unused.data(), unused.size()}}}, + {"too_small", first_desc, {{output_0.data(), output_0.size()}}}, + {"null_slice", first_desc, {{nullptr, first_desc.object_size}}}, + }; + auto read_results = backend_->BatchRead(requests); + ASSERT_EQ(read_results.size(), requests.size()); + ASSERT_TRUE(read_results[0].has_value()); + ASSERT_FALSE(read_results[1].has_value()); + EXPECT_EQ(read_results[1].error(), ErrorCode::INVALID_PARAMS); + ASSERT_FALSE(read_results[2].has_value()); + EXPECT_EQ(read_results[2].error(), ErrorCode::INVALID_PARAMS); + + EXPECT_EQ(std::memcmp(first_value.data(), output_0.data(), output_0.size()), + 0); + EXPECT_EQ(std::memcmp(first_value.data() + output_0.size(), output_1.data(), + output_1.size()), + 0); + for (size_t i = 0; i < unused.size(); ++i) { + EXPECT_EQ(unused.data()[i], 'U'); + } +} + +TEST_F(DfsBackendTest, BatchReadPreservesPerKeyErrors) { + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = tmp_->path(); + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = tmp_->path(); + distributed_config.fs_adapter_type = "posix"; + distributed_config.shard_count = 4; + distributed_config.shard_capacity = 64 * 1024 * 1024; + distributed_config.alignment = 4096; + + auto adapter = std::make_unique(); + auto* controlled_adapter = adapter.get(); + auto backend = std::make_unique( + file_config, distributed_config, std::move(adapter)); + ASSERT_TRUE(backend->Init().has_value()); + + AlignedBuffer write_buf(4096); + ASSERT_NE(write_buf.data(), nullptr); + write_buf.Fill('T'); + std::vector writes; + for (size_t i = 0; i < 4; ++i) { + writes.push_back({"key_" + std::to_string(i), + {ShardPath(0), i * 4096, 4096, 4096, 0}, + {{write_buf.data(), write_buf.size()}}}); + } + auto write_results = backend->BatchWrite(writes); + ASSERT_EQ(write_results.size(), writes.size()); + for (const auto& result : write_results) { + ASSERT_TRUE(result.has_value()); + } + + controlled_adapter->FailReadCall(2); + controlled_adapter->ShortReadCall(3); + AlignedBuffer out0(4096), out1(4096), out2(4096), out3(4096); + std::vector reads{ + {"ok_0", writes[0].descriptor, {{out0.data(), out0.size()}}}, + {"failed", writes[1].descriptor, {{out1.data(), out1.size()}}}, + {"short", writes[2].descriptor, {{out2.data(), out2.size()}}}, + {"ok_3", writes[3].descriptor, {{out3.data(), out3.size()}}}, + }; + auto read_results = backend->BatchRead(reads); + ASSERT_EQ(read_results.size(), reads.size()); + EXPECT_TRUE(read_results[0].has_value()); + ASSERT_FALSE(read_results[1].has_value()); + EXPECT_EQ(read_results[1].error(), ErrorCode::FILE_OPEN_FAIL); + ASSERT_FALSE(read_results[2].has_value()); + EXPECT_EQ(read_results[2].error(), ErrorCode::FILE_READ_FAIL); + EXPECT_TRUE(read_results[3].has_value()); + EXPECT_EQ(std::memcmp(write_buf.data(), out3.data(), write_buf.size()), 0); +} + +TEST_F(DfsBackendTest, RejectsInvalidDescriptorRangesBeforeIo) { + constexpr uint64_t kShardCapacity = 64 * 1024 * 1024; + + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = tmp_->path(); + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = tmp_->path(); + distributed_config.fs_adapter_type = "posix"; + distributed_config.shard_count = 4; + distributed_config.shard_capacity = kShardCapacity; + distributed_config.alignment = 4096; + + auto adapter = std::make_unique(); + auto* controlled_adapter = adapter.get(); + auto backend = std::make_unique( + file_config, distributed_config, std::move(adapter)); + ASSERT_TRUE(backend->Init().has_value()); + + AlignedBuffer small_buf(4096); + AlignedBuffer large_buf(8192); + ASSERT_NE(small_buf.data(), nullptr); + ASSERT_NE(large_buf.data(), nullptr); + + std::vector requests{ + {"past_capacity", + {ShardPath(0), kShardCapacity, 4096, 4096, 0}, + {{small_buf.data(), small_buf.size()}}}, + {"object_exceeds_allocation", + {ShardPath(0), 0, 8192, 4096, 0}, + {{large_buf.data(), large_buf.size()}}}, + {"overflow", + {ShardPath(0), std::numeric_limits::max() - 4095, 4096, 4096, + 0}, + {{small_buf.data(), small_buf.size()}}}, + {"unaligned_offset", + {ShardPath(0), 1, 4096, 4096, 0}, + {{small_buf.data(), small_buf.size()}}}, + }; + + auto write_results = backend->BatchWrite(requests); + ASSERT_EQ(write_results.size(), requests.size()); + for (const auto& result : write_results) { + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error(), ErrorCode::INVALID_PARAMS); + } + EXPECT_EQ(controlled_adapter->WriteCallCount(), 0); + + AlignedBuffer read_buf(8192); + ASSERT_NE(read_buf.data(), nullptr); + std::vector read_requests; + for (const auto& request : requests) { + read_requests.push_back({request.key, + request.descriptor, + {{read_buf.data(), read_buf.size()}}}); + } + auto read_results = backend->BatchRead(read_requests); + ASSERT_EQ(read_results.size(), read_requests.size()); + for (const auto& result : read_results) { + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error(), ErrorCode::INVALID_PARAMS); + } + EXPECT_EQ(controlled_adapter->ReadCallCount(), 0); +} + +TEST_F(DfsBackendTest, KeyOnlyStorageBackendOperationsAreNotSupported) { + std::unordered_map> offload_batch; + auto offload_result = backend_->BatchOffload( + offload_batch, + [](const std::vector&, + std::vector&) { return ErrorCode::OK; }); + ASSERT_FALSE(offload_result.has_value()); + EXPECT_EQ(offload_result.error(), ErrorCode::NOT_SUPPORTED); + + std::unordered_map load_batch; + auto load_result = backend_->BatchLoad(load_batch); + ASSERT_FALSE(load_result.has_value()); + EXPECT_EQ(load_result.error(), ErrorCode::NOT_SUPPORTED); + + auto exists_result = backend_->IsExist("key"); + ASSERT_FALSE(exists_result.has_value()); + EXPECT_EQ(exists_result.error(), ErrorCode::NOT_SUPPORTED); + auto enabled_result = backend_->IsEnableOffloading(); + ASSERT_TRUE(enabled_result.has_value()); + EXPECT_FALSE(*enabled_result); +} + +TEST_F(DfsBackendTest, BatchReadAndWriteAcceptUnalignedBuffers) { + constexpr size_t kObjectSize = 1234; + AlignedBuffer write_storage(kObjectSize + 1); + AlignedBuffer read_storage(kObjectSize + 3); + ASSERT_NE(write_storage.data(), nullptr); + ASSERT_NE(read_storage.data(), nullptr); + + char* write_ptr = write_storage.data() + 1; + char* read_ptr = read_storage.data() + 3; + ASSERT_NE(reinterpret_cast(write_ptr) % 4096, 0); + ASSERT_NE(reinterpret_cast(read_ptr) % 4096, 0); + for (size_t i = 0; i < kObjectSize; ++i) { + write_ptr[i] = static_cast('a' + (i % 26)); + } + + const DistributedFSDescriptor descriptor{ShardPath(0), 0, kObjectSize, 4096, + 0}; + auto write_results = backend_->BatchWrite( + {{"unaligned", descriptor, {{write_ptr, kObjectSize}}}}); + ASSERT_EQ(write_results.size(), 1); + ASSERT_TRUE(write_results[0].has_value()); + + auto read_results = backend_->BatchRead( + {{"unaligned", descriptor, {{read_ptr, kObjectSize}}}}); + ASSERT_EQ(read_results.size(), 1); + ASSERT_TRUE(read_results[0].has_value()); + EXPECT_EQ(std::memcmp(write_ptr, read_ptr, kObjectSize), 0); +} + +TEST_F(DfsBackendTest, MultipleKeysAcrossShards) { + AlignedBuffer key0(4096), key1(4096), key2(8192); + ASSERT_NE(key0.data(), nullptr); + ASSERT_NE(key1.data(), nullptr); + ASSERT_NE(key2.data(), nullptr); + key0.Fill('A'); + key1.Fill('B'); + key2.Fill('C'); + + const std::vector descriptors{ + {ShardPath(0), 0, key0.size(), key0.size(), 0}, + {ShardPath(0), 4096, key1.size(), key1.size(), 0}, + {ShardPath(1), 0, key2.size(), key2.size(), 1}, + }; + auto write_results = backend_->BatchWrite( + {{"key0", descriptors[0], {{key0.data(), key0.size()}}}, + {"key1", descriptors[1], {{key1.data(), key1.size()}}}, + {"key2", descriptors[2], {{key2.data(), key2.size()}}}}); + ASSERT_EQ(write_results.size(), 3); + EXPECT_TRUE(write_results[0].has_value()); + EXPECT_TRUE(write_results[1].has_value()); + EXPECT_TRUE(write_results[2].has_value()); + + AlignedBuffer out0(4096), out1(4096), out2(8192); + ASSERT_NE(out0.data(), nullptr); + ASSERT_NE(out1.data(), nullptr); + ASSERT_NE(out2.data(), nullptr); + auto read_results = backend_->BatchRead( + {{"key0", descriptors[0], {{out0.data(), out0.size()}}}, + {"key1", descriptors[1], {{out1.data(), out1.size()}}}, + {"key2", descriptors[2], {{out2.data(), out2.size()}}}}); + ASSERT_EQ(read_results.size(), 3); + EXPECT_TRUE(read_results[0].has_value()); + EXPECT_TRUE(read_results[1].has_value()); + EXPECT_TRUE(read_results[2].has_value()); + EXPECT_EQ(std::memcmp(key0.data(), out0.data(), key0.size()), 0); + EXPECT_EQ(std::memcmp(key1.data(), out1.data(), key1.size()), 0); + EXPECT_EQ(std::memcmp(key2.data(), out2.data(), key2.size()), 0); +} + +TEST_F(DfsBackendTest, LargeObject) { + constexpr size_t kLargeSize = 33 * 1024 * 1024; + AlignedBuffer write_buf(kLargeSize); + ASSERT_NE(write_buf.data(), nullptr); + write_buf.Fill('L'); + const DistributedFSDescriptor descriptor{ShardPath(2), 0, write_buf.size(), + write_buf.size(), 2}; + auto write_results = backend_->BatchWrite( + {{"large", descriptor, {{write_buf.data(), write_buf.size()}}}}); + ASSERT_EQ(write_results.size(), 1); + ASSERT_TRUE(write_results[0].has_value()); + + AlignedBuffer read_buf(kLargeSize); + ASSERT_NE(read_buf.data(), nullptr); + auto read_results = backend_->BatchRead( + {{"large", descriptor, {{read_buf.data(), read_buf.size()}}}}); + ASSERT_EQ(read_results.size(), 1); + ASSERT_TRUE(read_results[0].has_value()); + EXPECT_EQ(std::memcmp(write_buf.data(), read_buf.data(), kLargeSize), 0); +} + +TEST_F(DfsBackendTest, FailurePaths) { + PosixFsAdapter adapter; + ASSERT_TRUE(adapter.Init(tmp_->path()).has_value()); + char buf[64] = {}; + iovec iov{buf, sizeof(buf)}; + auto write = adapter.WriteAt(-1, &iov, 1, 0); + ASSERT_FALSE(write.has_value()); + EXPECT_EQ(write.error(), ErrorCode::INVALID_PARAMS); + + auto read = + adapter.ReadFile(tmp_->file("does_not_exist"), buf, sizeof(buf)); + ASSERT_FALSE(read.has_value()); + EXPECT_EQ(read.error(), ErrorCode::FILE_NOT_FOUND); +} + +TEST(DfsStorageFactoryTest, CreatesDistributedBackendWithPosixAdapter) { + EnvGuard env; + ConfigurePosixDfs(env); + TempDir tmp("dfs_factory"); + env.Set("MOONCAKE_DFS_ROOT_DIR", tmp.path().c_str()); + env.Set("MOONCAKE_DFS_SHARD_COUNT", "2"); + env.Set("MOONCAKE_DFS_SHARD_CAPACITY", "1048576"); + env.Set("MOONCAKE_DFS_ALIGNMENT", "4096"); + env.Set("MOONCAKE_DFS_SINGLE_TENANT", "true"); + + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = tmp.path(); + + auto backend = CreateStorageBackend(file_config); + ASSERT_TRUE(backend.has_value()); + ASSERT_NE(*backend, nullptr); + EXPECT_TRUE((*backend)->Init().has_value()); +} + +} // namespace mooncake::test diff --git a/mooncake-store/tests/dfs_sync_client_test.cpp b/mooncake-store/tests/dfs_sync_client_test.cpp new file mode 100644 index 0000000000..873cad93f2 --- /dev/null +++ b/mooncake-store/tests/dfs_sync_client_test.cpp @@ -0,0 +1,363 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "client_service.h" +#include "storage/distributed/dfs_global_allocator.h" +#include "storage/distributed/distributed_storage_backend.h" +#include "storage/distributed/posix_fs_adapter.h" +#include "test_server_helpers.h" +#include "utils.h" + +namespace mooncake::test { + +class FailingPosixFsAdapter : public PosixFsAdapter { + public: + int WriteCalls() const { return write_calls_.load(); } + void FailWriteCall(int call) { fail_write_call_.store(call); } + + tl::expected WriteAt(int fd, const iovec* iov, + int iovcnt, + int64_t offset) override { + const int call = ++write_calls_; + if (call == fail_write_call_.load()) { + return tl::make_unexpected(ErrorCode::FILE_WRITE_FAIL); + } + return PosixFsAdapter::WriteAt(fd, iov, iovcnt, offset); + } + + private: + std::atomic write_calls_{0}; + std::atomic fail_write_call_{-1}; +}; + +class DfsSyncClientTest : public ::testing::Test { + protected: + void SetUp() override { + root_ = (std::filesystem::temp_directory_path() / + ("dfs_sync_client_" + std::to_string(::getpid()) + "_" + + std::to_string(++next_root_))) + .string(); + std::filesystem::create_directories(root_); + + SetEnv("MOONCAKE_ENABLE_DFS", "1"); + SetEnv("MOONCAKE_DFS_FS_ADAPTER", "posix"); + SetEnv("MOONCAKE_DFS_ROOT_DIR", root_); + SetEnv("MOONCAKE_DFS_SHARD_COUNT", "2"); + SetEnv("MOONCAKE_DFS_SHARD_CAPACITY", "16777216"); + SetEnv("MOONCAKE_DFS_ALIGNMENT", "4096"); + SetEnv("MOONCAKE_DFS_EVICTION_ENABLED", "0"); + SetEnv("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "0"); + SetEnv("MOONCAKE_DFS_SINGLE_TENANT", "true"); + + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())); + writer_ = CreateClient("127.0.0.1:18101"); + provider_ = CreateClient("127.0.0.1:18102"); + ASSERT_NE(writer_, nullptr); + ASSERT_NE(provider_, nullptr); + + segment_size_ = 16 * 1024 * 1024; + segment_ = allocate_buffer_allocator_memory(segment_size_); + ASSERT_NE(segment_, nullptr); + ASSERT_TRUE(provider_->MountSegment(segment_, segment_size_, "tcp") + .has_value()); + + FileStorageConfig file_config; + file_config.storage_backend_type = StorageBackendType::kDistributed; + file_config.storage_filepath = root_; + + DistributedStorageConfig distributed_config; + distributed_config.fsdir = root_; + distributed_config.fs_adapter_type = "posix"; + distributed_config.shard_count = 2; + distributed_config.shard_capacity = 16 * 1024 * 1024; + distributed_config.alignment = 4096; + + auto adapter = std::make_unique(); + adapter_ = adapter.get(); + backend_ = std::make_shared( + file_config, distributed_config, std::move(adapter)); + ASSERT_TRUE(backend_->Init().has_value()); + writer_->SetDfsStorageBackend(backend_); + } + + void TearDown() override { + if (provider_ && segment_) { + (void)provider_->UnmountSegment(segment_, segment_size_); + } + writer_.reset(); + provider_.reset(); + backend_.reset(); + master_.Stop(); + if (segment_) { + std::free(segment_); + segment_ = nullptr; + } + RestoreEnv(); + std::error_code ec; + std::filesystem::remove_all(root_, ec); + } + + std::shared_ptr CreateClient(const std::string& hostname) { + auto client = Client::Create(hostname, "P2PHANDSHAKE", "tcp", + std::nullopt, master_.master_address()); + return client ? *client : nullptr; + } + + ReplicateConfig DfsConfig() const { + ReplicateConfig config; + config.replica_num = 1; + config.dfs_replica_num = 1; + return config; + } + + tl::expected QueryDfsOnly(const std::string& key) { + auto query = writer_->Query(key); + if (!query) { + return tl::make_unexpected(query.error()); + } + std::vector replicas; + for (const auto& replica : query->replicas) { + if (replica.is_dfs_replica()) { + replicas.push_back(replica); + } + } + if (replicas.empty()) { + return tl::make_unexpected(ErrorCode::INVALID_REPLICA); + } + return QueryResult(std::move(replicas), query->lease_timeout, + query->object_checksum); + } + + void ExpectDfsValue(const std::string& key, const std::string& expected) { + auto query = QueryDfsOnly(key); + ASSERT_TRUE(query.has_value()); + std::vector output(expected.size()); + const auto& descriptor = query->replicas[0].get_dfs_descriptor(); + auto results = backend_->BatchRead( + {{key, descriptor, {{output.data(), output.size()}}}}); + ASSERT_EQ(results.size(), 1); + ASSERT_TRUE(results[0].has_value()); + EXPECT_EQ(std::memcmp(output.data(), expected.data(), expected.size()), + 0); + } + + static std::vector> MakeSlices( + std::vector& values) { + std::vector> slices; + slices.reserve(values.size()); + for (auto& value : values) { + slices.push_back({Slice{value.data(), value.size()}}); + } + return slices; + } + + void SetEnv(const std::string& key, const std::string& value) { + const char* old_value = ::getenv(key.c_str()); + saved_env_.push_back({key, old_value + ? std::optional(old_value) + : std::nullopt}); + ::setenv(key.c_str(), value.c_str(), 1); + } + + void RestoreEnv() { + for (auto it = saved_env_.rbegin(); it != saved_env_.rend(); ++it) { + if (it->second) { + ::setenv(it->first.c_str(), it->second->c_str(), 1); + } else { + ::unsetenv(it->first.c_str()); + } + } + saved_env_.clear(); + } + + inline static std::atomic next_root_{0}; + testing::InProcMaster master_; + std::shared_ptr writer_; + std::shared_ptr provider_; + std::shared_ptr backend_; + FailingPosixFsAdapter* adapter_ = nullptr; + void* segment_ = nullptr; + size_t segment_size_ = 0; + std::string root_; + std::vector>> saved_env_; +}; + +TEST_F(DfsSyncClientTest, PutAndUpsertReturnAfterDfsWrite) { + const std::string key = "sync_put_upsert"; + std::string initial(4096, 'A'); + std::vector initial_slices{{initial.data(), initial.size()}}; + ASSERT_TRUE(writer_->Put(key, initial_slices, DfsConfig()).has_value()); + ExpectDfsValue(key, initial); + + std::string updated(4096, 'B'); + std::vector updated_slices{{updated.data(), updated.size()}}; + ASSERT_TRUE(writer_->Upsert(key, updated_slices, DfsConfig()).has_value()); + ExpectDfsValue(key, updated); +} + +TEST_F(DfsSyncClientTest, BatchPutAndBatchUpsertReturnAfterDfsWrite) { + std::vector keys{"sync_batch_0", "sync_batch_1"}; + std::vector initial{std::string(4096, 'C'), + std::string(4096, 'D')}; + auto initial_slices = MakeSlices(initial); + auto put_results = writer_->BatchPut(keys, initial_slices, DfsConfig()); + ASSERT_EQ(put_results.size(), keys.size()); + ASSERT_TRUE(put_results[0].has_value()); + ASSERT_TRUE(put_results[1].has_value()); + ExpectDfsValue(keys[0], initial[0]); + ExpectDfsValue(keys[1], initial[1]); + + std::vector updated{std::string(4096, 'E'), + std::string(4096, 'F')}; + auto updated_slices = MakeSlices(updated); + auto upsert_results = + writer_->BatchUpsert(keys, updated_slices, DfsConfig()); + ASSERT_EQ(upsert_results.size(), keys.size()); + ASSERT_TRUE(upsert_results[0].has_value()); + ASSERT_TRUE(upsert_results[1].has_value()); + ExpectDfsValue(keys[0], updated[0]); + ExpectDfsValue(keys[1], updated[1]); +} + +TEST_F(DfsSyncClientTest, PreferSameNodeBatchPutReturnsAfterDfsWrite) { + std::vector keys{"sync_same_node_0", "sync_same_node_1"}; + std::vector values{std::string(4096, 'J'), + std::string(4096, 'K')}; + auto slices = MakeSlices(values); + auto config = DfsConfig(); + config.prefer_alloc_in_same_node = true; + + auto results = writer_->BatchPut(keys, slices, config); + ASSERT_EQ(results.size(), keys.size()); + ASSERT_TRUE(results[0].has_value()); + ASSERT_TRUE(results[1].has_value()); + ExpectDfsValue(keys[0], values[0]); + ExpectDfsValue(keys[1], values[1]); +} + +TEST_F(DfsSyncClientTest, BatchDfsFailureOnlyRevokesFailedKey) { + std::vector keys{"sync_success", "sync_failure"}; + std::vector values{std::string(4096, 'G'), + std::string(4096, 'H')}; + auto slices = MakeSlices(values); + adapter_->FailWriteCall(adapter_->WriteCalls() + 2); + + auto results = writer_->BatchPut(keys, slices, DfsConfig()); + ASSERT_EQ(results.size(), keys.size()); + EXPECT_TRUE(results[0].has_value()); + ASSERT_FALSE(results[1].has_value()); + EXPECT_EQ(results[1].error(), ErrorCode::FILE_WRITE_FAIL); + ExpectDfsValue(keys[0], values[0]); + + auto failed_query = writer_->Query(keys[1]); + ASSERT_FALSE(failed_query.has_value()); + EXPECT_EQ(failed_query.error(), ErrorCode::OBJECT_NOT_FOUND); +} + +TEST_F(DfsSyncClientTest, MissingDfsBackendRevokesObject) { + auto client_without_backend = CreateClient("127.0.0.1:18103"); + ASSERT_NE(client_without_backend, nullptr); + std::string value(4096, 'I'); + std::vector slices{{value.data(), value.size()}}; + + auto result = + client_without_backend->Put("sync_no_backend", slices, DfsConfig()); + ASSERT_FALSE(result.has_value()); + EXPECT_EQ(result.error(), ErrorCode::DFS_SERVICE_UNAVAILABLE); + auto query = client_without_backend->Query("sync_no_backend"); + ASSERT_FALSE(query.has_value()); + EXPECT_EQ(query.error(), ErrorCode::OBJECT_NOT_FOUND); +} + +TEST_F(DfsSyncClientTest, GetReadsDfsIntoMultipleSlices) { + const std::string key = "dfs_multi_slice_get"; + std::string value(4096, 'M'); + std::vector write_slices{{value.data(), value.size()}}; + ASSERT_TRUE(writer_->Put(key, write_slices, DfsConfig()).has_value()); + + auto query = QueryDfsOnly(key); + ASSERT_TRUE(query.has_value()); + std::vector first(1024), second(3072); + std::vector read_slices{{first.data(), first.size()}, + {second.data(), second.size()}}; + ASSERT_TRUE(writer_->Get(key, *query, read_slices).has_value()); + EXPECT_EQ(std::memcmp(value.data(), first.data(), first.size()), 0); + EXPECT_EQ( + std::memcmp(value.data() + first.size(), second.data(), second.size()), + 0); +} + +TEST_F(DfsSyncClientTest, BatchGetUsesExplicitDfsDescriptors) { + std::vector keys{"dfs_batch_read_0", "dfs_batch_read_1"}; + std::vector values{std::string(4096, 'N'), + std::string(4096, 'O')}; + auto write_slices = MakeSlices(values); + auto put_results = writer_->BatchPut(keys, write_slices, DfsConfig()); + ASSERT_EQ(put_results.size(), keys.size()); + ASSERT_TRUE(put_results[0].has_value()); + ASSERT_TRUE(put_results[1].has_value()); + + std::vector queries; + queries.reserve(keys.size()); + for (const auto& key : keys) { + auto query = QueryDfsOnly(key); + ASSERT_TRUE(query.has_value()); + queries.push_back(*query); + } + + std::vector output{std::string(4096, '\0'), + std::string(4096, '\0')}; + std::unordered_map> read_slices; + for (size_t i = 0; i < keys.size(); ++i) { + read_slices[keys[i]] = {{output[i].data(), output[i].size()}}; + } + auto get_results = writer_->BatchGet(keys, queries, read_slices); + ASSERT_EQ(get_results.size(), keys.size()); + ASSERT_TRUE(get_results[0].has_value()); + ASSERT_TRUE(get_results[1].has_value()); + EXPECT_EQ(output[0], values[0]); + EXPECT_EQ(output[1], values[1]); +} + +TEST_F(DfsSyncClientTest, BatchGetVerifiesDfsChecksum) { + const char* checksum_enabled = std::getenv("MOONCAKE_STORE_CHECKSUM"); + if (checksum_enabled == nullptr || std::string(checksum_enabled) != "1") { + GTEST_SKIP() << "MOONCAKE_STORE_CHECKSUM is not enabled"; + } + + const std::string key = "dfs_batch_checksum"; + std::string value(4096, 'P'); + std::vector write_slices{{value.data(), value.size()}}; + ASSERT_TRUE(writer_->Put(key, write_slices, DfsConfig()).has_value()); + + auto query = QueryDfsOnly(key); + ASSERT_TRUE(query.has_value()); + ASSERT_TRUE(query->object_checksum.has_value()); + std::vector replicas = query->replicas; + std::vector queries; + queries.emplace_back( + std::move(replicas), query->lease_timeout, + std::optional(*query->object_checksum ^ uint64_t{1})); + + std::string output(value.size(), '\0'); + std::unordered_map> read_slices; + read_slices[key] = {{output.data(), output.size()}}; + auto results = writer_->BatchGet({key}, queries, read_slices); + ASSERT_EQ(results.size(), 1); + ASSERT_FALSE(results[0].has_value()); + EXPECT_EQ(results[0].error(), ErrorCode::CHECKSUM_MISMATCH); +} + +} // namespace mooncake::test diff --git a/mooncake-store/tests/master_scenario.cpp b/mooncake-store/tests/master_scenario.cpp index 9b06c867dd..142ab2248c 100644 --- a/mooncake-store/tests/master_scenario.cpp +++ b/mooncake-store/tests/master_scenario.cpp @@ -411,10 +411,8 @@ MasterScenario& MasterScenario::Then(TenantQuotaSpec tenant_quota) { } const auto matches = [&tenant_quota](const TenantQuotaSnapshot& value) { - return (!tenant_quota.used_bytes.has_value() || - value.used_bytes == *tenant_quota.used_bytes) && - (!tenant_quota.reserved_bytes.has_value() || - value.reserved_bytes == *tenant_quota.reserved_bytes); + return !tenant_quota.charged_bytes.has_value() || + value.charged_bytes == *tenant_quota.charged_bytes; }; const auto deadline = std::chrono::steady_clock::now() + tenant_quota.eventual_timeout; @@ -428,17 +426,11 @@ MasterScenario& MasterScenario::Then(TenantQuotaSpec tenant_quota) { } } - if (tenant_quota.used_bytes.has_value() && - snapshot->used_bytes != *tenant_quota.used_bytes) { - Fail("TenantQuota(" + tenant_quota.tenant + ") uses " + - std::to_string(snapshot->used_bytes) + "; expected " + - std::to_string(*tenant_quota.used_bytes)); - } - if (tenant_quota.reserved_bytes.has_value() && - snapshot->reserved_bytes != *tenant_quota.reserved_bytes) { - Fail("TenantQuota(" + tenant_quota.tenant + ") reserves " + - std::to_string(snapshot->reserved_bytes) + "; expected " + - std::to_string(*tenant_quota.reserved_bytes)); + if (tenant_quota.charged_bytes.has_value() && + snapshot->charged_bytes != *tenant_quota.charged_bytes) { + Fail("TenantQuota(" + tenant_quota.tenant + ") charges " + + std::to_string(snapshot->charged_bytes) + "; expected " + + std::to_string(*tenant_quota.charged_bytes)); } return *this; } diff --git a/mooncake-store/tests/master_scenario.h b/mooncake-store/tests/master_scenario.h index 9614d70496..c530695666 100644 --- a/mooncake-store/tests/master_scenario.h +++ b/mooncake-store/tests/master_scenario.h @@ -599,17 +599,11 @@ KeyCountSpec KeyCount(size_t value); struct TenantQuotaSpec { std::string tenant; - std::optional used_bytes{}; - std::optional reserved_bytes{}; + std::optional charged_bytes{}; std::chrono::milliseconds eventual_timeout{}; - TenantQuotaSpec& Uses(uint64_t value) { - used_bytes = value; - return *this; - } - - TenantQuotaSpec& Reserves(uint64_t value) { - reserved_bytes = value; + TenantQuotaSpec& Charges(uint64_t value) { + charged_bytes = value; return *this; } diff --git a/mooncake-store/tests/master_service_evict_scenario_test.cpp b/mooncake-store/tests/master_service_evict_scenario_test.cpp index bbc2e29095..d978071295 100644 --- a/mooncake-store/tests/master_service_evict_scenario_test.cpp +++ b/mooncake-store/tests/master_service_evict_scenario_test.cpp @@ -491,18 +491,18 @@ TEST_F(MasterServiceEvictScenarioTest, backend->BlockTxn(); scenario.When(EvictMemory(1.0)) .Then(Object("cold").DoesNotExist()) - .Then(TenantQuota(tenant).Uses(kObjectSize).Reserves(0)) + .Then(TenantQuota(tenant).Charges(kObjectSize)) .When(PutStart("before-durable", kObjectSize) .ForTenant(tenant) .ExpectError(ErrorCode::TENANT_QUOTA_EXCEEDED)); backend->AllowTxn(); ReadBatchEventually(storage, 3, batch); - scenario.Then(TenantQuota(tenant).Uses(0).Reserves(0).Eventually()) + scenario.Then(TenantQuota(tenant).Charges(0).Eventually()) .When(PutStart("after-durable", kObjectSize) .ForTenant(tenant) .ExpectReplicas(1)) - .Then(TenantQuota(tenant).Uses(0).Reserves(kObjectSize)); + .Then(TenantQuota(tenant).Charges(kObjectSize)); } TEST_F(MasterServiceEvictScenarioTest, @@ -558,8 +558,8 @@ TEST_F(MasterServiceEvictScenarioTest, .When(EvictMemory(0.5)) .Then(Object("same-key").ForTenant("tenant-a").DoesNotExist()) .Then(Object("same-key").ForTenant("tenant-b").IsReadable()) - .Then(TenantQuota("tenant-a").Uses(0).Reserves(0)) - .Then(TenantQuota("tenant-b").Uses(kObjectSize).Reserves(0)); + .Then(TenantQuota("tenant-a").Charges(0)) + .Then(TenantQuota("tenant-b").Charges(kObjectSize)); } TEST_F(MasterServiceEvictScenarioTest, @@ -582,10 +582,10 @@ TEST_F(MasterServiceEvictScenarioTest, .ExpectReplicas(1)) .Then(Object("tenant-a-old").ForTenant("tenant-a").DoesNotExist()) .Then(Object("tenant-b-object").ForTenant("tenant-b").IsReadable()) - .Then(TenantQuota("tenant-a").Uses(0).Reserves(kObjectSize)) + .Then(TenantQuota("tenant-a").Charges(kObjectSize)) .When(PutEnd("tenant-a-new").ForTenant("tenant-a")) - .Then(TenantQuota("tenant-a").Uses(kObjectSize).Reserves(0)) - .Then(TenantQuota("tenant-b").Uses(kObjectSize).Reserves(0)); + .Then(TenantQuota("tenant-a").Charges(kObjectSize)) + .Then(TenantQuota("tenant-b").Charges(kObjectSize)); } } // namespace diff --git a/mooncake-store/tests/master_service_test.cpp b/mooncake-store/tests/master_service_test.cpp index 02b072782a..e55264c4ff 100644 --- a/mooncake-store/tests/master_service_test.cpp +++ b/mooncake-store/tests/master_service_test.cpp @@ -978,6 +978,202 @@ TEST_F(MasterServiceTest, PutStartOnePlusOneAllowsSingleAllocatedReplica) { } #endif +TEST_F(MasterServiceTest, DfsPutEndAllAndUpsertTopologyAreAtomic) { + const auto dfs_root = (std::filesystem::temp_directory_path() / + ("master_dfs_sync_" + std::to_string(::getpid()))) + .string(); + std::filesystem::create_directories(dfs_root); + ScopedEnvVar enable_dfs("MOONCAKE_ENABLE_DFS", "1"); + ScopedEnvVar fs_adapter("MOONCAKE_DFS_FS_ADAPTER", "posix"); + ScopedEnvVar root_dir("MOONCAKE_DFS_ROOT_DIR", dfs_root.c_str()); + ScopedEnvVar shard_count("MOONCAKE_DFS_SHARD_COUNT", "1"); + ScopedEnvVar shard_capacity("MOONCAKE_DFS_SHARD_CAPACITY", "1048576"); + ScopedEnvVar alignment("MOONCAKE_DFS_ALIGNMENT", "4096"); + ScopedEnvVar eviction("MOONCAKE_DFS_EVICTION_ENABLED", "0"); + ScopedEnvVar deferred_free("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "0"); + ScopedEnvVar single_tenant("MOONCAKE_DFS_SINGLE_TENANT", "true"); + + { + MasterService service; + const auto context = PrepareSimpleSegment(service); + ReplicateConfig config; + config.replica_num = 1; + config.dfs_replica_num = 1; + + auto start = service.PutStart(context.client_id, "dfs_atomic", + TenantId::Default(), 4096, config); + ASSERT_TRUE(start.has_value()); + ASSERT_EQ(start->size(), 2); + ASSERT_TRUE(service + .PutEnd(context.client_id, "dfs_atomic", + TenantId::Default(), ReplicaType::ALL) + .has_value()); + + auto query = service.GetReplicaList("dfs_atomic", TenantId::Default()); + ASSERT_TRUE(query.has_value()); + ASSERT_EQ(query->replicas.size(), 2); + for (const auto& replica : query->replicas) { + EXPECT_EQ(replica.status, ReplicaStatus::COMPLETE); + } + + ReplicateConfig mismatched_config; + mismatched_config.replica_num = 1; + auto upsert = + service.UpsertStart(context.client_id, "dfs_atomic", + TenantId::Default(), 4096, mismatched_config); + ASSERT_FALSE(upsert.has_value()); + EXPECT_EQ(upsert.error(), ErrorCode::INVALID_PARAMS); + + query = service.GetReplicaList("dfs_atomic", TenantId::Default()); + ASSERT_TRUE(query.has_value()); + ASSERT_EQ(query->replicas.size(), 2); + for (const auto& replica : query->replicas) { + EXPECT_EQ(replica.status, ReplicaStatus::COMPLETE); + } + + auto revoke_start = service.PutStart(context.client_id, "dfs_revoke", + TenantId::Default(), 4096, config); + ASSERT_TRUE(revoke_start.has_value()); + ASSERT_TRUE(service + .PutRevoke(context.client_id, "dfs_revoke", + TenantId::Default(), ReplicaType::ALL) + .has_value()); + auto revoked = + service.GetReplicaList("dfs_revoke", TenantId::Default()); + ASSERT_FALSE(revoked.has_value()); + EXPECT_EQ(revoked.error(), ErrorCode::OBJECT_NOT_FOUND); + } + + { + MasterService service(MakeStrictTenantConfig({"default"})); + const auto context = PrepareSimpleSegment(service); + ReplicateConfig dfs_config; + dfs_config.replica_num = 1; + dfs_config.dfs_replica_num = 1; + + auto failed = service.PutStart(context.client_id, "dfs_quota_failure", + TenantId::Default(), + kStrictTenantQuotaBytes, dfs_config); + ASSERT_FALSE(failed.has_value()); + EXPECT_EQ(failed.error(), ErrorCode::NO_AVAILABLE_HANDLE); + + ReplicateConfig memory_config; + memory_config.replica_num = 1; + auto retry = service.PutStart( + context.client_id, "quota_after_dfs_failure", TenantId::Default(), + kStrictTenantQuotaBytes, memory_config); + ASSERT_TRUE(retry.has_value()) << toString(retry.error()); + ASSERT_TRUE(service + .PutRevoke(context.client_id, "quota_after_dfs_failure", + TenantId::Default(), ReplicaType::ALL) + .has_value()); + } + std::error_code ec; + std::filesystem::remove_all(dfs_root, ec); +} + +TEST_F(MasterServiceTest, DfsEvictionSplitsAcceptedAndRejectedCandidates) { + auto run_case = [&](const std::string& case_name, + const std::vector& leased_indexes, + const std::vector>& + expected_replica_counts, + bool evict_memory_first = false) { + const auto dfs_root = (std::filesystem::temp_directory_path() / + ("master_dfs_evict_" + + std::to_string(::getpid()) + "_" + case_name)) + .string(); + std::filesystem::create_directories(dfs_root); + ScopedEnvVar enable_dfs("MOONCAKE_ENABLE_DFS", "1"); + ScopedEnvVar fs_adapter("MOONCAKE_DFS_FS_ADAPTER", "posix"); + ScopedEnvVar root_dir("MOONCAKE_DFS_ROOT_DIR", dfs_root.c_str()); + ScopedEnvVar shard_count("MOONCAKE_DFS_SHARD_COUNT", "1"); + ScopedEnvVar shard_capacity("MOONCAKE_DFS_SHARD_CAPACITY", "32768"); + ScopedEnvVar alignment("MOONCAKE_DFS_ALIGNMENT", "4096"); + // Keep the background path disabled so the test drives one exact + // transaction through the public test hook. + ScopedEnvVar eviction("MOONCAKE_DFS_EVICTION_ENABLED", "0"); + ScopedEnvVar high_watermark("MOONCAKE_DFS_EVICTION_HIGH_WATERMARK", + "0.9"); + ScopedEnvVar low_watermark("MOONCAKE_DFS_EVICTION_LOW_WATERMARK", + "0.7"); + ScopedEnvVar deferred_free("MOONCAKE_DFS_DEFERRED_FREE_SECONDS", "0"); + ScopedEnvVar single_tenant("MOONCAKE_DFS_SINGLE_TENANT", "true"); + + { + MasterService service; + const auto context = PrepareSimpleSegment(service); + ReplicateConfig config; + config.replica_num = 1; + config.dfs_replica_num = 1; + + std::vector keys; + for (int i = 0; i < 4; ++i) { + keys.push_back("dfs_evict_" + std::to_string(i)); + auto start = service.PutStart(context.client_id, keys.back(), + TenantId::Default(), 100, config); + ASSERT_TRUE(start.has_value()) << "allocation " << i; + ASSERT_TRUE(service + .PutEnd(context.client_id, keys.back(), + TenantId::Default(), ReplicaType::ALL) + .has_value()); + } + + for (const size_t index : leased_indexes) { + ASSERT_LT(index, keys.size()); + ASSERT_TRUE( + service.GetReplicaList(keys[index], TenantId::Default()) + .has_value()); + } + + if (evict_memory_first) { + service.RunBatchEvictForTesting(1.0, 1.0); + } + service.RunDfsEvictionForTesting(); + + for (size_t i = 0; i < keys.size(); ++i) { + auto result = + service.GetReplicaList(keys[i], TenantId::Default()); + if (!expected_replica_counts[i].has_value()) { + ASSERT_FALSE(result.has_value()) << "key=" << keys[i]; + EXPECT_EQ(result.error(), ErrorCode::OBJECT_NOT_FOUND) + << "key=" << keys[i]; + continue; + } + ASSERT_TRUE(result.has_value()); + EXPECT_EQ(result->replicas.size(), *expected_replica_counts[i]) + << "key=" << keys[i]; + } + + if (evict_memory_first) { + auto reclaimed = + service.PutStart(context.client_id, "dfs_evict_reclaimed", + TenantId::Default(), 100, config); + ASSERT_TRUE(reclaimed.has_value()) << reclaimed.error(); + ASSERT_TRUE( + service + .PutRevoke(context.client_id, "dfs_evict_reclaimed", + TenantId::Default(), ReplicaType::ALL) + .has_value()); + } + } + + std::error_code ec; + std::filesystem::remove_all(dfs_root, ec); + }; + + run_case("commit", {}, {1, 1, 2, 2}); + run_case("reject", {0, 1, 2, 3}, {2, 2, 2, 2}); + // k1 shares the first prepared batch with k0. Rejecting k1 must not roll + // back k0, and the same high-watermark trigger must continue to k2 so the + // shard reaches its low watermark. + run_case("mixed", {1}, {1, 2, 1, 2}); + // Memory eviction may leave DFS as the only remaining replica. DFS + // eviction must still reclaim those allocations and erase metadata for + // objects whose final replica was removed. + run_case("last_replica", {}, + {std::nullopt, std::nullopt, size_t{1}, size_t{1}}, true); +} + TEST_F(MasterServiceTest, PutStartGroupIdsValidation) { std::unique_ptr service_(new MasterService()); [[maybe_unused]] const auto context = PrepareSimpleSegment(*service_); diff --git a/mooncake-store/tests/replica_selection_test.cpp b/mooncake-store/tests/replica_selection_test.cpp index 388cede25a..af5600bad7 100644 --- a/mooncake-store/tests/replica_selection_test.cpp +++ b/mooncake-store/tests/replica_selection_test.cpp @@ -67,6 +67,14 @@ Replica::Descriptor MakeLocalDisk(const std::string& endpoint) { return d; } +Replica::Descriptor MakeDfs(const std::string& path) { + Replica::Descriptor d; + d.id = 0; + d.descriptor_variant = DistributedFSDescriptor{path, 0, 1024, 4096, 0}; + d.status = ReplicaStatus::COMPLETE; + return d; +} + // A test fixture that guarantees scoring state is reset between tests, since // the enable flag / injected scorer are process-wide. class ReplicaSelectionTest : public ::testing::Test { @@ -280,6 +288,28 @@ TEST_F(ReplicaSelectionTest, DiskIsLastCompleteFallback) { EXPECT_TRUE(sel->is_disk_replica()); } +TEST_F(ReplicaSelectionTest, DfsPrecedesDisk) { + std::unordered_set local; + std::vector reps = { + MakeDisk("/remote/object"), + MakeDfs("/dfs/object"), + }; + const auto* sel = SelectBestReplica(reps, local); + ASSERT_NE(sel, nullptr); + EXPECT_TRUE(sel->is_dfs_replica()); +} + +TEST_F(ReplicaSelectionTest, LocalDiskPrecedesDfs) { + std::unordered_set local; + std::vector reps = { + MakeDfs("/dfs/object"), + MakeLocalDisk("nodeA"), + }; + const auto* sel = SelectBestReplica(reps, local); + ASSERT_NE(sel, nullptr); + EXPECT_TRUE(sel->is_local_disk_replica()); +} + TEST_F(ReplicaSelectionTest, NoCompleteReplicaReturnsNull) { std::unordered_set local; std::vector reps = { From 32f98620cc1e0e10c2276815f78e8bf9a097bba0 Mon Sep 17 00:00:00 2001 From: Xuanlin_Mi <179552030+Primary33@users.noreply.github.com> Date: Thu, 13 Aug 2026 10:45:59 +0800 Subject: [PATCH 046/483] [TENT] Port GDS and Ascend transport reliability updates (#3280) * [TENT] Port TE transport reliability updates * [TENT] Drop incomplete TCP CUDA context port * [TENT] Report GDS partial progress while pending * [TENT] Defer GDS terminal status until IO drains --- .../ascend/ascend_direct_transport.h | 4 +- .../tent/transport/gds/gds_transport.h | 37 ++- .../ascend/ascend_direct_transport.cpp | 23 +- .../tent/src/transport/gds/gds_transport.cpp | 292 +++++++++++++++--- .../tent/tests/CMakeLists.txt | 13 + .../tent/tests/gds_transport_status_test.cpp | 233 ++++++++++++++ 6 files changed, 554 insertions(+), 48 deletions(-) create mode 100644 mooncake-transfer-engine/tent/tests/gds_transport_status_test.cpp diff --git a/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h b/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h index 8b55c24e83..e4be7fd06e 100644 --- a/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h +++ b/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h @@ -82,6 +82,8 @@ class AscendDirectTransport : public Transport { void disconnect(const std::string &remote_hixl, int32_t timeout_in_millis); + void forgetConnectedSegment(const std::string &remote_hixl); + void startTransfer(SegmentID target_id, Request::OpCode opcode, const std::vector &tasks); @@ -131,4 +133,4 @@ class AscendDirectTransport : public Transport { } // namespace tent } // namespace mooncake -#endif // ASCEND_DIRECT_TRANSPORT_H_ \ No newline at end of file +#endif // ASCEND_DIRECT_TRANSPORT_H_ diff --git a/mooncake-transfer-engine/tent/include/tent/transport/gds/gds_transport.h b/mooncake-transfer-engine/tent/include/tent/transport/gds/gds_transport.h index d63c42bfe6..2b48f399cf 100644 --- a/mooncake-transfer-engine/tent/include/tent/transport/gds/gds_transport.h +++ b/mooncake-transfer-engine/tent/include/tent/transport/gds/gds_transport.h @@ -34,13 +34,17 @@ namespace mooncake { namespace tent { class GdsFileContext; +class GdsTransportTestPeer; struct IOParamRange { size_t base = 0; size_t count = 0; - size_t complete_count = 0; size_t transferred_bytes = 0; TransferStatusEnum status = TransferStatusEnum::PENDING; + // Public status stays PENDING until every slice in this range is + // physically terminal. Preserve the first observed failure while sibling + // slices drain so cancellation cannot mask the original result. + TransferStatusEnum known_failure = TransferStatusEnum::PENDING; }; // Wrapper for reusable CUfileBatchHandle_t @@ -55,7 +59,13 @@ struct GdsSubBatch : public Transport::SubBatch { BatchHandle* batch_handle; // Pointer to reusable handle from pool std::vector io_param_ranges; std::vector io_params; + // Scratch buffer populated by cuFileBatchIOGetStatus(). std::vector io_events; + // Stable per-slice status cache indexed by the cookie carried in each IO. + std::vector cached_events; + bool reusable = true; + bool cancel_requested = false; + std::mutex status_mutex; virtual size_t size() const { return io_param_ranges.size(); } }; @@ -90,6 +100,24 @@ class GdsTransport : public Transport { virtual const char* getName() const { return "gds"; } private: + friend class GdsTransportTestPeer; + + static TransferStatus aggregateTransferStatus( + const std::vector& events, size_t base, size_t count, + bool& all_terminal); + + Status updateBatchStatus(GdsSubBatch* batch); + + Status cancelBatch(GdsSubBatch* batch); + + static bool allBatchIOsTerminal(const GdsSubBatch* batch); + + static bool isTerminalFailure(TransferStatusEnum status); + + static void destroySubBatch(GdsSubBatch* batch); + + void cleanupQuarantinedBatches(); + std::string getGdsFilePath(SegmentID handle); GdsFileContext* findFileContext(SegmentID target_id); @@ -115,8 +143,13 @@ class GdsTransport : public Transport { // Track all allocated sub-batches to clean up on uninstall std::vector allocated_batches_; std::mutex allocated_batches_lock_; + + // Failed batches may still be referenced by cuFile after cancellation. + // Keep their handle and IO storage alive until every slice is terminal. + std::vector quarantined_batches_; + std::mutex quarantined_batches_lock_; }; } // namespace tent } // namespace mooncake -#endif // GDS_TRANSPORT_H_ \ No newline at end of file +#endif // GDS_TRANSPORT_H_ diff --git a/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp b/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp index dd93a23254..598c39170c 100644 --- a/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp @@ -360,8 +360,13 @@ void AscendDirectTransport::startTransfer( LOG(ERROR) << "Failed to transfer to: " << remote_hixl << ", status: " << hixl_ret << ", errmsg: " << aclGetRecentErrMsg(); - // disconnect to remote when transfer fail - disconnect(remote_hixl, 10); + // AutoConnect tears down failed routes inside HIXL. Calling + // Disconnect again is redundant; only discard our local bookkeeping. + if (auto_connect_) { + forgetConnectedSegment(remote_hixl); + } else { + disconnect(remote_hixl, 10); + } for (auto &task : tasks) { task->status_word = TransferStatusEnum::FAILED; } @@ -430,7 +435,13 @@ Status AscendDirectTransport::getTransferStatus(SubBatchRef batch, int task_id, LOG(ERROR) << "Get transfer status failed, ret: " << hixlTransferStatusToString(xfer_status) << ", errmsg: " << aclGetRecentErrMsg(); - disconnect(task.remote_hixl, 10); + // HIXL DisconnectOnError has already handled AutoConnect failures. + // Explicit application timeouts still use disconnect() above. + if (auto_connect_) { + forgetConnectedSegment(task.remote_hixl); + } else { + disconnect(task.remote_hixl, 10); + } task.status_word = TransferStatusEnum::FAILED; } if (task.batch_size > 1) { @@ -474,6 +485,12 @@ void AscendDirectTransport::disconnect(const std::string &remote_hixl, } } +void AscendDirectTransport::forgetConnectedSegment( + const std::string &remote_hixl) { + std::lock_guard lock(connection_mutex_); + connected_segments_.erase(remote_hixl); +} + Status AscendDirectTransport::addMemoryBuffer(BufferDesc &desc, const MemoryOptions &options) { desc.transports.push_back(TransportType::AscendDirect); diff --git a/mooncake-transfer-engine/tent/src/transport/gds/gds_transport.cpp b/mooncake-transfer-engine/tent/src/transport/gds/gds_transport.cpp index 6980d98c15..9961b192cd 100644 --- a/mooncake-transfer-engine/tent/src/transport/gds/gds_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/gds/gds_transport.cpp @@ -87,6 +87,158 @@ TransferStatusEnum parseTransferStatus(CUfileStatus_t status) { } } +bool isTerminalCuFileStatus(CUfileStatus_t status) { + return status != CUFILE_WAITING && status != CUFILE_PENDING; +} + +TransferStatus GdsTransport::aggregateTransferStatus( + const std::vector& events, size_t base, size_t count, + bool& all_terminal) { + all_terminal = true; + TransferStatus result{COMPLETED, 0}; + if (count == 0 || base > events.size() || count > events.size() - base) { + result.s = INVALID; + return result; + } + + // Use a fixed precedence so completion order does not change the result. + int failure_priority = 0; + for (size_t i = base; i < base + count; ++i) { + const auto& event = events[i]; + const auto slice_status = parseTransferStatus(event.status); + switch (slice_status) { + case PENDING: + all_terminal = false; + break; + case COMPLETED: + if (event.ret > 0) { + result.transferred_bytes += static_cast(event.ret); + } + break; + case INVALID: + if (failure_priority < 1) { + result.s = INVALID; + failure_priority = 1; + } + break; + case CANCELED: + if (failure_priority < 2) { + result.s = CANCELED; + failure_priority = 2; + } + break; + case TIMEOUT: + if (failure_priority < 3) { + result.s = TIMEOUT; + failure_priority = 3; + } + break; + case FAILED: + result.s = FAILED; + failure_priority = 4; + break; + case INITIAL: + all_terminal = false; + break; + } + } + + if (!all_terminal && failure_priority == 0) { + result.s = PENDING; + } + return result; +} + +Status GdsTransport::updateBatchStatus(GdsSubBatch* batch) { + if (batch->io_events.size() < batch->io_params.size() || + batch->cached_events.size() < batch->io_params.size()) { + return Status::InternalError( + "GDS completion buffers are smaller than the submitted IO " + "set" LOC_MARK); + } + + unsigned num_events = static_cast(batch->io_params.size()); + if (num_events == 0) return Status::OK(); + + auto result = + cuFileBatchIOGetStatus(batch->batch_handle->handle, 0, &num_events, + batch->io_events.data(), nullptr); + if (result.err != CU_FILE_SUCCESS) { + return Status::InternalError( + std::string("Failed to get GDS batch status: Code ") + + std::to_string(result.err) + LOC_MARK); + } + + for (size_t index = 0; index < num_events; ++index) { + const auto& event = batch->io_events[index]; + const auto cookie = reinterpret_cast(event.cookie); + if (cookie == 0 || cookie > batch->cached_events.size()) { + LOG(ERROR) << "Invalid GDS batch IO cookie: " << cookie; + continue; + } + + auto& cached_event = batch->cached_events[cookie - 1]; + if (!isTerminalCuFileStatus(cached_event.status) || + isTerminalCuFileStatus(event.status)) { + cached_event = event; + } + } + return Status::OK(); +} + +Status GdsTransport::cancelBatch(GdsSubBatch* batch) { + auto result = cuFileBatchIOCancel(batch->batch_handle->handle); + if (result.err != CU_FILE_SUCCESS) { + return Status::InternalError( + std::string("Failed to cancel GDS batch IO: Code ") + + std::to_string(result.err) + LOC_MARK); + } + return Status::OK(); +} + +bool GdsTransport::allBatchIOsTerminal(const GdsSubBatch* batch) { + if (batch->cached_events.size() < batch->io_params.size()) return false; + return std::all_of(batch->cached_events.begin(), + batch->cached_events.begin() + batch->io_params.size(), + [](const CUfileIOEvents_t& event) { + return isTerminalCuFileStatus(event.status); + }); +} + +bool GdsTransport::isTerminalFailure(TransferStatusEnum status) { + return status == INVALID || status == CANCELED || status == TIMEOUT || + status == FAILED; +} + +void GdsTransport::destroySubBatch(GdsSubBatch* batch) { + if (batch->batch_handle) { + cuFileBatchIODestroy(batch->batch_handle->handle); + delete batch->batch_handle; + batch->batch_handle = nullptr; + } + Slab::Get().deallocate(batch); +} + +void GdsTransport::cleanupQuarantinedBatches() { + std::lock_guard lock(quarantined_batches_lock_); + auto it = quarantined_batches_.begin(); + while (it != quarantined_batches_.end()) { + auto* batch = *it; + bool ready_to_destroy = false; + { + std::lock_guard status_lock(batch->status_mutex); + auto status = updateBatchStatus(batch); + ready_to_destroy = status.ok() && allBatchIOsTerminal(batch); + } + if (!ready_to_destroy) { + ++it; + continue; + } + destroySubBatch(batch); + it = quarantined_batches_.erase(it); + } +} + GdsTransport::GdsTransport() : installed_(false) { static std::once_flag g_once_flag; auto fork_init = []() { cuFileDriverOpen(); }; @@ -121,16 +273,19 @@ Status GdsTransport::uninstall() { { std::lock_guard lock(allocated_batches_lock_); for (auto* gds_batch : allocated_batches_) { - // Destroy the batch handle (don't return to pool since we're - // shutting down) - cuFileBatchIODestroy(gds_batch->batch_handle->handle); - delete gds_batch->batch_handle; - // Deallocate the sub-batch - Slab::Get().deallocate(gds_batch); + destroySubBatch(gds_batch); } allocated_batches_.clear(); } + { + std::lock_guard lock(quarantined_batches_lock_); + for (auto* gds_batch : quarantined_batches_) { + destroySubBatch(gds_batch); + } + quarantined_batches_.clear(); + } + // Clean up all handles in the pool std::lock_guard lock(handle_pool_lock_); for (auto* batch_handle : handle_pool_) { @@ -146,6 +301,8 @@ Status GdsTransport::uninstall() { } Status GdsTransport::allocateSubBatch(SubBatchRef& batch, size_t max_size) { + cleanupQuarantinedBatches(); + auto gds_batch = Slab::Get().allocate(); if (!gds_batch) return Status::InternalError("Unable to allocate GDS sub-batch"); @@ -189,6 +346,10 @@ Status GdsTransport::allocateSubBatch(SubBatchRef& batch, size_t max_size) { gds_batch->io_params.clear(); gds_batch->io_params.reserve(io_batch_depth_); gds_batch->io_param_ranges.clear(); + gds_batch->cached_events.clear(); + gds_batch->cached_events.reserve(io_batch_depth_); + gds_batch->reusable = true; + gds_batch->cancel_requested = false; // Track this batch for cleanup on uninstall { @@ -215,18 +376,28 @@ Status GdsTransport::freeSubBatch(SubBatchRef& batch) { } } - // Return the handle to pool for reuse (avoid expensive - // cuFileBatchIODestroy) Note: Caller should ensure all IOs are completed - // (via getTransferStatus) before calling freeSubBatch, as cuFile may still - // access io_params otherwise + // A normal batch is returned to the handle pool only after callers have + // observed terminal status for every task. A failed batch with active + // slices retains both its handle and parameter storage until cuFile no + // longer references them. + bool reusable = false; { - std::lock_guard lock(handle_pool_lock_); - handle_pool_.push_back(gds_batch->batch_handle); + std::lock_guard lock(gds_batch->status_mutex); + reusable = gds_batch->reusable; + } + if (reusable) { + { + std::lock_guard lock(handle_pool_lock_); + handle_pool_.push_back(gds_batch->batch_handle); + } + gds_batch->batch_handle = nullptr; + Slab::Get().deallocate(gds_batch); + } else { + std::lock_guard lock(quarantined_batches_lock_); + quarantined_batches_.push_back(gds_batch); } - - // Deallocate the GdsSubBatch (each allocation gets a fresh one) - Slab::Get().deallocate(gds_batch); batch = nullptr; + cleanupQuarantinedBatches(); return Status::OK(); } @@ -279,24 +450,30 @@ Status GdsTransport::submitTransferTasks( GdsFileContext* context = findFileContext(request.target_id); if (!context || !context->ready()) return Status::InvalidArgument("Invalid remote segment" LOC_MARK); - size_t task_id = gds_batch->io_param_ranges.size(); IOParamRange range; range.base = gds_batch->io_params.size(); for (size_t offset = 0; offset < request.length; offset += kMaxSliceSize) { size_t length = std::min(kMaxSliceSize, request.length - offset); + const size_t slice_id = gds_batch->io_params.size(); CUfileIOParams_t params; params.mode = CUFILE_BATCH; params.opcode = (request.opcode == Request::READ) ? CUFILE_READ : CUFILE_WRITE; - params.cookie = - reinterpret_cast(static_cast(task_id)); + // Use a one-based slice index so every completion can be cached + // independently, including the first slice (cookie 0 is avoided). + params.cookie = reinterpret_cast( + static_cast(slice_id + 1)); params.u.batch.devPtr_base = request.source; params.u.batch.devPtr_offset = offset; params.u.batch.file_offset = request.target_offset + offset; params.u.batch.size = length; params.fh = context->getHandle(); gds_batch->io_params.push_back(params); + CUfileIOEvents_t cached_event{}; + cached_event.cookie = params.cookie; + cached_event.status = CUFILE_PENDING; + gds_batch->cached_events.push_back(cached_event); range.count++; } gds_batch->io_param_ranges.push_back(range); @@ -315,39 +492,70 @@ Status GdsTransport::submitTransferTasks( Status GdsTransport::getTransferStatus(SubBatchRef batch, int task_id, TransferStatus& status) { auto gds_batch = dynamic_cast(batch); + if (!gds_batch) + return Status::InvalidArgument("Invalid GDS sub-batch" LOC_MARK); unsigned num_tasks = gds_batch->io_param_ranges.size(); if (task_id < 0 || task_id >= (int)num_tasks) return Status::InvalidArgument("Invalid task ID"); - unsigned num_events = static_cast(gds_batch->io_params.size()); - auto result = - cuFileBatchIOGetStatus(gds_batch->batch_handle->handle, 0, &num_events, - gds_batch->io_events.data(), nullptr); - if (result.err != CU_FILE_SUCCESS) - return Status::InternalError( - std::string("Failed to get GDS batch status: Code ") + - std::to_string(result.err) + LOC_MARK); + std::lock_guard lock(gds_batch->status_mutex); + auto& range = gds_batch->io_param_ranges[task_id]; + if (range.status != PENDING) { + status = TransferStatus{range.status, range.transferred_bytes}; + return Status::OK(); + } - for (size_t index = 0; index < num_events; ++index) { - auto& event = gds_batch->io_events[index]; - auto event_task_id = reinterpret_cast(event.cookie); - if (event_task_id >= gds_batch->io_param_ranges.size()) { - LOG(ERROR) << "Invalid GDS batch IO cookie: " << event_task_id; - continue; + auto update_status = updateBatchStatus(gds_batch); + if (!update_status.ok()) { + // The runtime may release the sub-batch after a polling error. Keep + // cuFile-owned storage out of the reusable pool unless we observed a + // terminal event for every submitted slice. + gds_batch->reusable = false; + if (!gds_batch->cancel_requested) { + gds_batch->cancel_requested = true; + auto cancel_status = cancelBatch(gds_batch); + if (!cancel_status.ok()) { + LOG(WARNING) << cancel_status.ToString(); + } } + return update_status; + } + bool all_terminal = false; + auto task_status = aggregateTransferStatus( + gds_batch->cached_events, range.base, range.count, all_terminal); + if (range.known_failure == PENDING && isTerminalFailure(task_status.s)) { + range.known_failure = task_status.s; + } - auto& range = gds_batch->io_param_ranges[event_task_id]; - auto s = parseTransferStatus(event.status); - if (s == COMPLETED) { - range.complete_count++; - range.transferred_bytes += event.ret; - } else if (s != PENDING) { - range.status = s; + if (range.known_failure != PENDING && !all_terminal) { + if (!gds_batch->cancel_requested) { + gds_batch->cancel_requested = true; + auto cancel_status = cancelBatch(gds_batch); + if (!cancel_status.ok()) { + LOG(WARNING) << cancel_status.ToString(); + } + } + + // Cancellation is best effort. Poll once more, but do not publish a + // terminal status while cuFile may still access the user buffer. + auto repoll_status = updateBatchStatus(gds_batch); + if (!repoll_status.ok()) { + LOG(WARNING) << repoll_status.ToString(); + } + task_status = aggregateTransferStatus( + gds_batch->cached_events, range.base, range.count, all_terminal); + if (!all_terminal) { + gds_batch->reusable = false; } } - auto& range = gds_batch->io_param_ranges[task_id]; - if (range.complete_count == range.count) { - range.status = COMPLETED; + // Expose partial progress while the task is still pending, and keep the + // reported byte count monotonic across repeated polls. + range.transferred_bytes = + std::max(range.transferred_bytes, task_status.transferred_bytes); + if (all_terminal) { + range.status = range.known_failure != PENDING ? range.known_failure + : task_status.s; + gds_batch->reusable = allBatchIOsTerminal(gds_batch); } status = TransferStatus{range.status, range.transferred_bytes}; return Status::OK(); diff --git a/mooncake-transfer-engine/tent/tests/CMakeLists.txt b/mooncake-transfer-engine/tent/tests/CMakeLists.txt index 23ad86c835..5ee112896f 100644 --- a/mooncake-transfer-engine/tent/tests/CMakeLists.txt +++ b/mooncake-transfer-engine/tent/tests/CMakeLists.txt @@ -204,6 +204,19 @@ if(USE_HIP) add_test(NAME tent_rocm_platform_test COMMAND tent_rocm_platform_test) endif() +if(TARGET tent_xport_gds) + add_executable(tent_gds_transport_status_test gds_transport_status_test.cpp) + target_link_libraries(tent_gds_transport_status_test PRIVATE gtest gtest_main + tent_link_group) + target_link_options( + tent_gds_transport_status_test PRIVATE "-Wl,--wrap=cuFileDriverOpen" + "-Wl,--wrap=cuFileBatchIOGetStatus" "-Wl,--wrap=cuFileBatchIOCancel") + target_include_directories(tent_gds_transport_status_test + PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}/../include) + add_test(NAME tent_gds_transport_status_test + COMMAND tent_gds_transport_status_test) +endif() + if(USE_SUNRISE) add_executable(tent_sunrise_link_transport_test sunrise_link_transport_test.cpp) diff --git a/mooncake-transfer-engine/tent/tests/gds_transport_status_test.cpp b/mooncake-transfer-engine/tent/tests/gds_transport_status_test.cpp new file mode 100644 index 0000000000..a28df72b4c --- /dev/null +++ b/mooncake-transfer-engine/tent/tests/gds_transport_status_test.cpp @@ -0,0 +1,233 @@ +// Copyright 2026 KVCache.AI +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include + +#include +#include +#include + +#include "tent/transport/gds/gds_transport.h" + +namespace { + +std::vector reported_events; +size_t cancel_call_count = 0; + +} // namespace + +extern "C" CUfileError_t __wrap_cuFileDriverOpen() { + CUfileError_t result{}; + result.err = CU_FILE_SUCCESS; + return result; +} + +extern "C" CUfileError_t __wrap_cuFileBatchIOGetStatus(CUfileBatchHandle_t, + unsigned, + unsigned* num_events, + CUfileIOEvents_t* events, + timespec*) { + const auto count = std::min(*num_events, reported_events.size()); + std::copy_n(reported_events.begin(), count, events); + *num_events = static_cast(count); + + CUfileError_t result{}; + result.err = CU_FILE_SUCCESS; + return result; +} + +extern "C" CUfileError_t __wrap_cuFileBatchIOCancel(CUfileBatchHandle_t) { + ++cancel_call_count; + CUfileError_t result{}; + result.err = CU_FILE_SUCCESS; + return result; +} + +namespace mooncake { +namespace tent { + +class GdsTransportTestPeer { + public: + static TransferStatus aggregate(const std::vector& events, + bool& all_terminal) { + return GdsTransport::aggregateTransferStatus(events, 0, events.size(), + all_terminal); + } +}; + +namespace { + +CUfileIOEvents_t makeEvent(CUfileStatus_t status, int64_t bytes = 0, + size_t slice_id = 0) { + CUfileIOEvents_t event{}; + event.status = status; + event.ret = bytes; + event.cookie = reinterpret_cast(slice_id + 1); + return event; +} + +TEST(GdsTransportStatusTest, DetectsFailureWhileSiblingIsActive) { + bool all_terminal = true; + auto status = GdsTransportTestPeer::aggregate( + {makeEvent(CUFILE_FAILED), makeEvent(CUFILE_WAITING)}, all_terminal); + + EXPECT_EQ(status.s, FAILED); + EXPECT_FALSE(all_terminal); +} + +TEST(GdsTransportStatusTest, ReportsPendingWhenNoFailureIsKnown) { + bool all_terminal = true; + auto status = GdsTransportTestPeer::aggregate( + {makeEvent(CUFILE_COMPLETE, 1024), makeEvent(CUFILE_PENDING)}, + all_terminal); + + EXPECT_EQ(status.s, PENDING); + EXPECT_EQ(status.transferred_bytes, 1024); + EXPECT_FALSE(all_terminal); +} + +TEST(GdsTransportStatusTest, PublicStatusReportsCompletedBytesWhilePending) { + GdsTransport transport; + GdsSubBatch batch; + BatchHandle batch_handle{}; + batch.batch_handle = &batch_handle; + batch.io_params.resize(2); + batch.io_events.resize(2); + batch.cached_events = {makeEvent(CUFILE_PENDING, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + batch.io_param_ranges.push_back(IOParamRange{0, 2, 0, PENDING, PENDING}); + reported_events = {makeEvent(CUFILE_COMPLETE, 1024, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + + TransferStatus status{INITIAL, 0}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, PENDING); + EXPECT_EQ(status.transferred_bytes, 1024); + EXPECT_EQ(batch.io_param_ranges[0].transferred_bytes, 1024); +} + +TEST(GdsTransportStatusTest, KeepsFailurePendingUntilEverySliceIsTerminal) { + GdsTransport transport; + GdsSubBatch batch; + BatchHandle batch_handle{}; + batch.batch_handle = &batch_handle; + batch.io_params.resize(2); + batch.io_events.resize(2); + batch.cached_events = {makeEvent(CUFILE_PENDING, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + batch.io_param_ranges.push_back(IOParamRange{0, 2, 0, PENDING, PENDING}); + reported_events = {makeEvent(CUFILE_FAILED, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + cancel_call_count = 0; + + TransferStatus status{INITIAL, 0}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, PENDING); + EXPECT_EQ(batch.io_param_ranges[0].known_failure, FAILED); + EXPECT_FALSE(batch.reusable); + EXPECT_EQ(cancel_call_count, 1); + + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, PENDING); + EXPECT_EQ(cancel_call_count, 1); + + reported_events = {makeEvent(CUFILE_FAILED, 0, 0), + makeEvent(CUFILE_CANCELED, 0, 1)}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, FAILED); + EXPECT_TRUE(batch.reusable); + EXPECT_EQ(cancel_call_count, 1); +} + +TEST(GdsTransportStatusTest, CancellationDoesNotMaskOriginalFailure) { + GdsTransport transport; + GdsSubBatch batch; + BatchHandle batch_handle{}; + batch.batch_handle = &batch_handle; + batch.io_params.resize(2); + batch.io_events.resize(2); + batch.cached_events = {makeEvent(CUFILE_PENDING, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + batch.io_param_ranges.push_back(IOParamRange{0, 2, 0, PENDING, PENDING}); + reported_events = {makeEvent(CUFILE_INVALID, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + cancel_call_count = 0; + + TransferStatus status{INITIAL, 0}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, PENDING); + EXPECT_EQ(batch.io_param_ranges[0].known_failure, INVALID); + + reported_events = {makeEvent(CUFILE_INVALID, 0, 0), + makeEvent(CUFILE_CANCELED, 0, 1)}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, INVALID); + EXPECT_EQ(cancel_call_count, 1); +} + +TEST(GdsTransportStatusTest, ReusesHandleOnlyAfterWholeBatchIsTerminal) { + GdsTransport transport; + GdsSubBatch batch; + BatchHandle batch_handle{}; + batch.batch_handle = &batch_handle; + batch.io_params.resize(2); + batch.io_events.resize(2); + batch.cached_events = {makeEvent(CUFILE_PENDING, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + batch.io_param_ranges = {IOParamRange{0, 1, 0, PENDING, PENDING}, + IOParamRange{1, 1, 0, PENDING, PENDING}}; + reported_events = {makeEvent(CUFILE_FAILED, 0, 0), + makeEvent(CUFILE_PENDING, 0, 1)}; + cancel_call_count = 0; + + TransferStatus status{INITIAL, 0}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 0, status).ok()); + EXPECT_EQ(status.s, FAILED); + EXPECT_FALSE(batch.reusable); + EXPECT_EQ(cancel_call_count, 0); + + reported_events = {makeEvent(CUFILE_FAILED, 0, 0), + makeEvent(CUFILE_COMPLETE, 1024, 1)}; + ASSERT_TRUE(transport.getTransferStatus(&batch, 1, status).ok()); + EXPECT_EQ(status.s, COMPLETED); + EXPECT_TRUE(batch.reusable); +} + +TEST(GdsTransportStatusTest, AggregatesCompletedBytes) { + bool all_terminal = false; + auto status = GdsTransportTestPeer::aggregate( + {makeEvent(CUFILE_COMPLETE, 1024), makeEvent(CUFILE_COMPLETE, 2048)}, + all_terminal); + + EXPECT_EQ(status.s, COMPLETED); + EXPECT_EQ(status.transferred_bytes, 3072); + EXPECT_TRUE(all_terminal); +} + +TEST(GdsTransportStatusTest, FailurePrecedenceIsCompletionOrderIndependent) { + bool all_terminal = false; + auto first = GdsTransportTestPeer::aggregate( + {makeEvent(CUFILE_CANCELED), makeEvent(CUFILE_FAILED)}, all_terminal); + EXPECT_EQ(first.s, FAILED); + EXPECT_TRUE(all_terminal); + + auto second = GdsTransportTestPeer::aggregate( + {makeEvent(CUFILE_FAILED), makeEvent(CUFILE_CANCELED)}, all_terminal); + EXPECT_EQ(second.s, FAILED); + EXPECT_TRUE(all_terminal); +} + +} // namespace +} // namespace tent +} // namespace mooncake From 6e40517616c335d788e5d92e67bd40e308ddff1a Mon Sep 17 00:00:00 2001 From: ykwd Date: Thu, 13 Aug 2026 12:20:00 +0800 Subject: [PATCH 047/483] [Docs] Reorganize design docs and move tebench to performance (#3376) * move store-related design to the store's sub-pages; delete an unused index.md; change design doc order * Move tebench to performance section * Fix broken link * Update the summary for tebench docs --- docs/source/conf.py | 8 ++++ .../mooncake-store-deployment-guide.md | 2 +- .../conductor-architecture-design.md | 2 +- docs/source/design/index.md | 37 ------------------- docs/source/design/{ => store}/engram.md | 0 .../design/{ => store}/mooncake-store.md | 18 +++++---- .../ssd-free-ratio-first-allocation.md | 0 docs/source/design/{ => store}/ssd-offload.md | 2 +- .../{ => store}/unified-parallel-tensor-io.md | 0 docs/source/getting_started/quick-start.md | 2 +- docs/source/index.md | 15 +++----- docs/source/performance/mooncake/index.md | 4 +- .../tent => performance/mooncake}/tebench.md | 0 13 files changed, 31 insertions(+), 59 deletions(-) delete mode 100644 docs/source/design/index.md rename docs/source/design/{ => store}/engram.md (100%) rename docs/source/design/{ => store}/mooncake-store.md (98%) rename docs/source/design/{ => store}/ssd-free-ratio-first-allocation.md (100%) rename docs/source/design/{ => store}/ssd-offload.md (99%) rename docs/source/design/{ => store}/unified-parallel-tensor-io.md (100%) rename docs/source/{design/tent => performance/mooncake}/tebench.md (100%) diff --git a/docs/source/conf.py b/docs/source/conf.py index c10fad2189..f339396447 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -255,6 +255,14 @@ def linkcode_resolve(domain, info): # Preserve published URLs when documentation is reorganized. Redirect targets # are relative to the generated location of each legacy page. redirects = { + "design/mooncake-store": "store/mooncake-store.html", + "design/ssd-offload": "store/ssd-offload.html", + "design/ssd-free-ratio-first-allocation": + "store/ssd-free-ratio-first-allocation.html", + "design/engram": "store/engram.html", + "design/unified-parallel-tensor-io": + "store/unified-parallel-tensor-io.html", + "design/tent/tebench": "../../performance/mooncake/tebench.html", "deployment/ssd-offload": "ssd/ssd-offload.html", "deployment/nvmf-ssd-deployment-guide": "ssd/nvmf-ssd-deployment-guide.html", diff --git a/docs/source/deployment/mooncake-store-deployment-guide.md b/docs/source/deployment/mooncake-store-deployment-guide.md index 2c8f099e85..56811c7b2e 100644 --- a/docs/source/deployment/mooncake-store-deployment-guide.md +++ b/docs/source/deployment/mooncake-store-deployment-guide.md @@ -13,7 +13,7 @@ This guide covers minimal deployment, and operational tuning of Mooncake Store. **Metadata Service**: A separate service (etcd, Redis, or HTTP) used by the Transfer Engine for peer discovery and configuration. The master's embedded HTTP metadata server can replace an external etcd/Redis for simple deployments. We also provide a P2P handshake mechanism (`P2PHANDSHAKE`) that enables decentralized metadata management by storing metadata locally on each node, eliminating the need for a centralized service — this is the simplest metadata handshake method and the recommended starting point (see [Quick Start](#quick-start)). -For a detailed design discussion, see the [Mooncake Store Design](../design/mooncake-store.md). +For a detailed design discussion, see the [Mooncake Store Design](../design/store/mooncake-store.md). --- diff --git a/docs/source/design/conductor/conductor-architecture-design.md b/docs/source/design/conductor/conductor-architecture-design.md index 115156b5c4..d9980169c1 100644 --- a/docs/source/design/conductor/conductor-architecture-design.md +++ b/docs/source/design/conductor/conductor-architecture-design.md @@ -1,4 +1,4 @@ -# Mooncake Conductor Architecture +# Mooncake Conductor ## Overview diff --git a/docs/source/design/index.md b/docs/source/design/index.md deleted file mode 100644 index 507f5cbd9d..0000000000 --- a/docs/source/design/index.md +++ /dev/null @@ -1,37 +0,0 @@ ---- -orphan: true ---- - -# Design Documents - -Architecture and implementation details for Mooncake's storage, transfer, and -distributed execution components. - -## Core Architecture - -| Document | Description | -|----------|-------------| -| [Mooncake Architecture](architecture) | KVCache-centric disaggregated serving architecture. | -| [Mooncake Store](mooncake-store) | Distributed object and KV cache storage design. | -| [Transfer Engine](transfer-engine/index) | High-performance data movement architecture and transports. | -| [P2P Store](p2p-store) | Peer-to-peer checkpoint and object transfer design. | - -## Serving and Cache Systems - -| Document | Description | -|----------|-------------| -| [HiCache](hicache-design) | Hierarchical KV cache design. | -| [Engram](engram) | Distributed serving and cache architecture. | -| [Unified Parallel Tensor I/O](unified-parallel-tensor-io) | Parallel tensor storage and transfer model. | -| [SSD Offload](ssd-offload) | SSD-backed cache hierarchy design. | -| [SSD Free-Ratio-First Allocation](ssd-free-ratio-first-allocation) | Capacity-aware replica placement strategy. | - -## Distributed Execution and Routing - -| Document | Description | -|----------|-------------| -| [Mooncake PG](mooncake-backend-pg) | Elastic PyTorch process group. | -| [Mooncake EP](mooncake-ep) | Expert-parallel communication and recovery. | -| [TENT](tent/overview) | Next-generation transfer engine design. | -| [TENT Benchmark](tent/tebench) | TENT benchmark framework and methodology. | -| [Conductor](conductor/conductor-architecture-design) | Cache-aware request routing architecture. | diff --git a/docs/source/design/engram.md b/docs/source/design/store/engram.md similarity index 100% rename from docs/source/design/engram.md rename to docs/source/design/store/engram.md diff --git a/docs/source/design/mooncake-store.md b/docs/source/design/store/mooncake-store.md similarity index 98% rename from docs/source/design/mooncake-store.md rename to docs/source/design/store/mooncake-store.md index 98b6f1138a..4c9665821d 100644 --- a/docs/source/design/mooncake-store.md +++ b/docs/source/design/store/mooncake-store.md @@ -21,7 +21,7 @@ Key features of Mooncake Store include: ## Architecture -![architecture](../image/mooncake-store-preview.png) +![architecture](../../image/mooncake-store-preview.png) As shown in the figure above, there are two key components in Mooncake Store: **Master Service** and **Client**. @@ -67,7 +67,7 @@ The `Client` class provides the primary interface for Mooncake Store operations: | `BatchReplicaClear` | Batch clear replicas on specific segments | | `QueryByRegex` / `RemoveByRegex` | Query or delete objects matching a regex | -For full API signatures, parameter details, and usage examples, see the [Mooncake Store C++ API Reference](../api-reference/cpp/mooncake-store.md). +For full API signatures, parameter details, and usage examples, see the [Mooncake Store C++ API Reference](../../api-reference/cpp/mooncake-store.md). ## Master Service @@ -95,7 +95,7 @@ To reduce cache warm-up time after a master restart, the Master Service supports The Master Service can optionally enforce strict multi-tenant memory quota admission. This feature is disabled by default. When `enable_multi_tenants=false`, request tenant IDs are ignored for object placement, all objects use the `default` namespace, and tenant quota management requests return `UNAVAILABLE_IN_CURRENT_MODE`. -See [Multi-Tenant Deployment](../deployment/multi-tenancy.md) for configuration details. +See [Multi-Tenant Deployment](../../deployment/multi-tenancy.md) for configuration details. Effective quota is recomputed from the current registered memory capacity: @@ -473,7 +473,7 @@ Mooncake Store provides two concrete implementations of `BufferAllocatorBase`: **OffsetBufferAllocator (default and recommended)**: This allocator is derived from [OffsetAllocator](https://github.com/sebbbi/OffsetAllocator), which uses a custom bin-based allocation strategy that supports fast hard realtime `O(1)` offset allocation with minimal fragmentation. Mooncake Store optimizes this allocator based on the specific memory usage characteristics of LLM inference workloads, thereby enhancing memory utilization in LLM scenarios. -For measured utilization and allocation latency across LLM-style workloads, see [Allocator Performance](../performance/mooncake/allocator-benchmark-result.md). +For measured utilization and allocation latency across LLM-style workloads, see [Allocator Performance](../../performance/mooncake/allocator-benchmark-result.md). **CachelibBufferAllocator (deprecated)**: This allocator leverages Facebook's [CacheLib](https://github.com/facebook/CacheLib) to manage memory using a slab-based allocation strategy. It provides efficient memory allocation with good fragmentation resistance and is well-suited for high-performance scenarios. However, in our modified version, it does not handle workloads with highly variable object sizes effectively, so it is currently marked as deprecated. @@ -567,7 +567,7 @@ Valid values are: `random` (default), `free_ratio_first`, `ssd_free_ratio_first` **Use `local_first`** when inference workers and Mooncake Store memory segments are colocated and you want writes to prefer the writer's host before falling back to other hosts. For this strategy to work correctly, all writer and store processes on the same physical or logical host must use the same stable, globally unique host part in `local_hostname`. -For benchmark data comparing `random` and `free_ratio_first` across segment counts, replica counts, and skewed capacities, see [AllocationStrategy Performance](../performance/mooncake/allocation-strategy-benchmark-result.md). +For benchmark data comparing `random` and `free_ratio_first` across segment counts, replica counts, and skewed capacities, see [AllocationStrategy Performance](../../performance/mooncake/allocation-strategy-benchmark-result.md). #### Strategy Details @@ -786,7 +786,7 @@ This integration is **experimental** and incomplete; see the plugin page for det The descriptor-based DFS data plane can use the native HF3FS USRBIO API instead of POSIX I/O. Select it with `MOONCAKE_DFS_FS_ADAPTER=hf3fs`; the legacy `--root_fs_dir` option does not enable this path, and there is no automatic -fallback to POSIX. See the [HF3FS USRBIO adapter guide](../getting_started/plugin-usage/3FS-USRBIO-Plugin.md) +fallback to POSIX. See the [HF3FS USRBIO adapter guide](../../getting_started/plugin-usage/3FS-USRBIO-Plugin.md) for build prerequisites and configuration. ## Builtin Metadata Server @@ -818,7 +818,7 @@ To start the master service with the HTTP metadata server enabled: When enabled, the HTTP metadata server will start automatically and provide metadata services for the Mooncake Store cluster. This eliminates the need for an external etcd deployment, simplifying the setup process for development and testing environments. Note that the HTTP metadata server is designed for single-node deployments and does not provide the high availability features that etcd offers. For production environments requiring high availability, etcd is still the recommended choice. -For detailed guidance on monitoring master metrics, Prometheus endpoints, and health checks, see the [Observability guide](../getting_started/observability.md). +For detailed guidance on monitoring master metrics, Prometheus endpoints, and health checks, see the [Observability guide](../../getting_started/observability.md). ## Mooncake Store Python API @@ -842,4 +842,8 @@ When to bump the version: :maxdepth: 1 ssd-offload +unified-parallel-tensor-io +ssd-free-ratio-first-allocation +engram + ::: diff --git a/docs/source/design/ssd-free-ratio-first-allocation.md b/docs/source/design/store/ssd-free-ratio-first-allocation.md similarity index 100% rename from docs/source/design/ssd-free-ratio-first-allocation.md rename to docs/source/design/store/ssd-free-ratio-first-allocation.md diff --git a/docs/source/design/ssd-offload.md b/docs/source/design/store/ssd-offload.md similarity index 99% rename from docs/source/design/ssd-offload.md rename to docs/source/design/store/ssd-offload.md index 71cd369ded..a80a2c7c30 100644 --- a/docs/source/design/ssd-offload.md +++ b/docs/source/design/store/ssd-offload.md @@ -6,7 +6,7 @@ Mooncake Store supports offloading KV cache objects from distributed memory to l SSD offload is implemented as a background subsystem within the **real client** process. It is transparent to the application: a `Put` that would otherwise be evicted from memory is persisted to disk, and a `Get` that finds no memory replica automatically falls back to reading from SSD. -For multi-turn conversation benchmark results, see [Mooncake SSD Offload Benchmark](../performance/mooncake/ssd-offload-benchmark-results.md). +For multi-turn conversation benchmark results, see [Mooncake SSD Offload Benchmark](../../performance/mooncake/ssd-offload-benchmark-results.md). --- diff --git a/docs/source/design/unified-parallel-tensor-io.md b/docs/source/design/store/unified-parallel-tensor-io.md similarity index 100% rename from docs/source/design/unified-parallel-tensor-io.md rename to docs/source/design/store/unified-parallel-tensor-io.md diff --git a/docs/source/getting_started/quick-start.md b/docs/source/getting_started/quick-start.md index 570aa37b2c..69df2d5a47 100644 --- a/docs/source/getting_started/quick-start.md +++ b/docs/source/getting_started/quick-start.md @@ -125,4 +125,4 @@ allocation strategies, SSD offload, and runtime tuning, continue to the [Mooncake Store Deployment & Tuning Guide](../deployment/mooncake-store-deployment-guide.md). For API details, see the [Mooncake Store Python API](../api-reference/python/mooncake-store.md) -and [Mooncake Store design](../design/mooncake-store.md). +and [Mooncake Store design](../design/store/mooncake-store.md). diff --git a/docs/source/index.md b/docs/source/index.md index 4f980bcc30..9bf417894c 100644 --- a/docs/source/index.md +++ b/docs/source/index.md @@ -124,19 +124,14 @@ performance/vllm/index :maxdepth: 1 design/architecture -design/mooncake-store -design/p2p-store -design/mooncake-backend-pg -design/mooncake-ep design/transfer-engine/index -design/hicache-design -design/engram -design/unified-parallel-tensor-io design/tent/overview -design/tent/tebench +design/store/mooncake-store +design/mooncake-backend-pg +design/mooncake-ep +design/p2p-store design/conductor/conductor-architecture-design -design/ssd-offload -design/ssd-free-ratio-first-allocation +design/hicache-design ::: % API Documentation diff --git a/docs/source/performance/mooncake/index.md b/docs/source/performance/mooncake/index.md index d511245f60..6fd7559b26 100644 --- a/docs/source/performance/mooncake/index.md +++ b/docs/source/performance/mooncake/index.md @@ -1,6 +1,6 @@ # Mooncake Performance -Benchmarks evaluating Mooncake Store's core storage, allocation, and cache hierarchy behavior. +This section collects Mooncake performance evaluations and benchmark results for Mooncake core components. | Document | Area | Key Findings | |----------|------|---------------| @@ -8,6 +8,7 @@ Benchmarks evaluating Mooncake Store's core storage, allocation, and cache hiera | [Allocator Benchmark](allocator-benchmark-result) | Segment allocation | The optimized OffsetAllocator significantly improves utilization for uniform-size LLM KV cache allocation patterns. | | [Allocation Strategy Benchmark](allocation-strategy-benchmark-result) | Allocation routing | Compares random and free-ratio-first allocation across segments, replicas, skewed capacity, and DSA-style KV+indexer workloads. | | [SSD Offload Benchmark](ssd-offload-benchmark-results) | Cache hierarchy | SSD offload extends the KV cache hierarchy with NVMe, reducing the performance cliff after DRAM cache capacity is exhausted in long multi-turn conversations. | +| [Guide for tebench](tebench) | Transfer Engine | End-to-end bandwidth and latency benchmarking for classic TE and TENT backends across block size, batch size, and concurrency. | :::{toctree} :maxdepth: 1 @@ -17,4 +18,5 @@ storage-benchmark allocator-benchmark-result allocation-strategy-benchmark-result ssd-offload-benchmark-results +tebench ::: diff --git a/docs/source/design/tent/tebench.md b/docs/source/performance/mooncake/tebench.md similarity index 100% rename from docs/source/design/tent/tebench.md rename to docs/source/performance/mooncake/tebench.md From 90c5ce09aad2535044dab1f39b3bae3daa0c80a1 Mon Sep 17 00:00:00 2001 From: Xun Sun Date: Thu, 13 Aug 2026 14:04:49 +0800 Subject: [PATCH 048/483] [EP] decouple native EP from torch C++ APIs (#2883) Co-authored-by: yuechen-sys --- CMakeLists.txt | 26 +- mooncake-ep/BuildEpExt.cmake | 110 ----- mooncake-ep/CMakeLists.txt | 34 +- mooncake-ep/benchmarks/legacy_buffer_perf.cpp | 199 ++++++++ mooncake-ep/benchmarks/legacy_buffer_perf.py | 123 +++++ .../elastic/mooncake_ep_elastic_buffer.h | 74 +-- .../elastic/mooncake_ep_elastic_compiled.cuh | 5 +- .../elastic/mooncake_ep_elastic_exception.cuh | 2 +- .../elastic/mooncake_ep_elastic_launch.cuh | 3 +- .../elastic/mooncake_ep_elastic_ptx.cuh | 4 +- mooncake-ep/include/mooncake_ep_api.cuh | 2 +- mooncake-ep/include/mooncake_ep_buffer.h | 60 +-- mooncake-ep/include/mooncake_ep_configs.cuh | 10 +- mooncake-ep/include/mooncake_ep_device.h | 38 +- mooncake-ep/include/mooncake_ep_event.h | 55 ++- mooncake-ep/setup.py | 138 ------ mooncake-ep/src/CMakeLists.txt | 97 +++- mooncake-ep/src/ep_py.cpp | 79 +-- mooncake-ep/src/mooncake_ep_buffer.cpp | 312 ++++++------ .../src/mooncake_ep_elastic_buffer.cpp | 463 ++++++------------ mooncake-ep/src/mooncake_ep_kernel.cu | 52 +- mooncake-ep/tests/test_ep_grid.py | 8 +- mooncake-integration/CMakeLists.txt | 1 + mooncake-pg/src/CMakeLists.txt | 1 + mooncake-pg/torch/src/mooncake_backend.cpp | 23 +- mooncake-wheel/mooncake/ep.py | 9 +- .../mooncake/mooncake_elastic_buffer.py | 313 +++++++++--- mooncake-wheel/mooncake/mooncake_ep_buffer.py | 251 +++++++--- mooncake-wheel/tests/ep_test_utils.py | 3 +- mooncake-wheel/tests/test_mooncake_ep.py | 12 +- 30 files changed, 1357 insertions(+), 1150 deletions(-) delete mode 100644 mooncake-ep/BuildEpExt.cmake create mode 100644 mooncake-ep/benchmarks/legacy_buffer_perf.cpp create mode 100644 mooncake-ep/benchmarks/legacy_buffer_perf.py delete mode 100644 mooncake-ep/setup.py diff --git a/CMakeLists.txt b/CMakeLists.txt index a82f333c77..29db7903bf 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -121,12 +121,15 @@ if (WITH_EP) mooncake-pg/include $) else () - message(STATUS "WITH_EP enabled: building Mooncake EP and PG Python extensions") + message(STATUS "WITH_EP enabled: building Mooncake EP natively and PG via setup.py") if(USE_CUDA) find_package(CUDAToolkit REQUIRED) message(STATUS "Detected CUDA version: ${CUDAToolkit_VERSION}") endif() + add_subdirectory(mooncake-ep) + include_directories(mooncake-ep/include) + # EP_TORCH_VERSIONS: semicolon-separated list of PyTorch versions to build for. # Can be set via -DEP_TORCH_VERSIONS="2.9.1;2.8.0" or the EP_TORCH_VERSIONS env var. # Empty means build with the currently-installed torch. @@ -156,25 +159,6 @@ if (WITH_EP) string(REPLACE ";" "|" _ep_torch_versions_pipe "${EP_TORCH_VERSIONS}") string(REPLACE ";" "|" _torch_cuda_arch_list_pipe "${TORCH_CUDA_ARCH_LIST}") - add_custom_target(mooncake_ep_ext ALL - COMMAND ${CMAKE_COMMAND} -E make_directory "${EP_PG_STAGING_DIR}" - COMMAND ${CMAKE_COMMAND} - "-DSOURCE_DIR=${CMAKE_CURRENT_SOURCE_DIR}/mooncake-ep" - "-DEP_CUDA_MAJOR=${CUDAToolkit_VERSION_MAJOR}" - "-DEP_CUDA_MINOR=${CUDAToolkit_VERSION_MINOR}" - "-DEP_TORCH_VERSIONS=${_ep_torch_versions_pipe}" - "-DTORCH_CUDA_ARCH_LIST=${_torch_cuda_arch_list_pipe}" - "-DSTAGING_DIR=${EP_PG_STAGING_DIR}" - "-DENGINE_SO_PATH=$" - "-DPython3_EXECUTABLE=${Python3_EXECUTABLE}" - "-DEP_USE_MUSA=$,1,0>" - "-DEP_USE_MACA=$,1,0>" - -P "${CMAKE_CURRENT_SOURCE_DIR}/mooncake-ep/BuildEpExt.cmake" - COMMENT "Building Mooncake EP Python extension(s)" - DEPENDS engine - VERBATIM - ) - add_custom_target(mooncake_pg_ext ALL COMMAND ${CMAKE_COMMAND} -E make_directory "${EP_PG_STAGING_DIR}" COMMAND ${CMAKE_COMMAND} @@ -190,7 +174,7 @@ if (WITH_EP) "-DEP_USE_MACA=$,1,0>" -P "${CMAKE_CURRENT_SOURCE_DIR}/mooncake-pg/torch/BuildPgExt.cmake" COMMENT "Building Mooncake PG Python extension(s)" - DEPENDS mooncake_pg mooncake_pg_device mooncake_ep_ext + DEPENDS mooncake_pg mooncake_pg_device VERBATIM ) endif () diff --git a/mooncake-ep/BuildEpExt.cmake b/mooncake-ep/BuildEpExt.cmake deleted file mode 100644 index 4a5a661ed4..0000000000 --- a/mooncake-ep/BuildEpExt.cmake +++ /dev/null @@ -1,110 +0,0 @@ -# BuildEpExt.cmake - Build the Mooncake EP Python extension. -# -# Invoked at build time via cmake -P from the root CMakeLists.txt when -# WITH_EP=ON. Variables are passed with -D from the custom target: -# -# SOURCE_DIR - mooncake-ep source directory -# EP_CUDA_MAJOR - CUDA major version (integer) -# EP_TORCH_VERSIONS - pipe-separated (|) PyTorch versions to build for -# (empty = use the currently-installed torch) -# TORCH_CUDA_ARCH_LIST - pipe-separated CUDA arch list forwarded to torch -# STAGING_DIR - destination directory for the built .so files -# ENGINE_SO_PATH - absolute path to the built engine.cpython-XYZ.so -# EP_USE_MUSA - set to "1" when building for MUSA (MTLink path) -# EP_USE_MACA - set to "1" when building for MACA (MTLink path) - -cmake_minimum_required(VERSION 3.16) - -# Include common build utilities. -include("${SOURCE_DIR}/../mooncake-common/SetupPyTorchEnv.cmake") - -# Restore pipe-separated strings back to CMake semicolon-separated lists. -if(EP_TORCH_VERSIONS) - string(REPLACE "|" ";" EP_TORCH_VERSIONS "${EP_TORCH_VERSIONS}") -endif() -if(TORCH_CUDA_ARCH_LIST) - string(REPLACE "|" ";" TORCH_CUDA_ARCH_LIST "${TORCH_CUDA_ARCH_LIST}") -endif() - -# --------------------------------------------------------------------------- -# 1. Set up the build environment. -# --------------------------------------------------------------------------- -# Clear jobserver variables so that sub-processes started by setup.py do not -# try to connect to the parent ninja's jobserver pipe FDs, which are not -# inherited and cause: "ninja: error: Could not initialize jobserver: Invalid -# file descriptors". -set(ENV{MAKEFLAGS} "") -set(ENV{MFLAGS} "") -set(ENV{TORCH_CUDA_ARCH_LIST} "${TORCH_CUDA_ARCH_LIST}") -if(EP_USE_MUSA) - set(ENV{MOONCAKE_EP_USE_MUSA} "1") -else() - unset(ENV{MOONCAKE_EP_USE_MUSA}) -endif() -if(EP_USE_MACA) - set(ENV{MOONCAKE_EP_USE_MACA} "1") - if(DEFINED ENV{MACA_PATH}) - set(ENV{MACA_HOME} "$ENV{MACA_PATH}") - elseif(DEFINED ENV{MACA_HOME}) - set(ENV{MACA_PATH} "$ENV{MACA_HOME}") - endif() -else() - unset(ENV{MOONCAKE_EP_USE_MACA}) -endif() - -# --------------------------------------------------------------------------- -# 2. Ensure engine.so exists in mooncake-wheel/mooncake/ for setup.py linking. -# --------------------------------------------------------------------------- -# setup.py links against -l:engine.so in ../mooncake-wheel/mooncake/. -# During the make phase only the versioned engine.cpython-XYZ.so exists in -# the build tree; create a bare engine.so symlink so the linker can find it. -set(_wheel_mooncake_dir "${SOURCE_DIR}/../mooncake-wheel/mooncake") -set(_engine_symlink "${_wheel_mooncake_dir}/engine.so") -if(ENGINE_SO_PATH AND NOT EXISTS "${_engine_symlink}") - message(STATUS "[EP] Creating engine.so symlink -> ${ENGINE_SO_PATH}") - execute_process( - COMMAND ${CMAKE_COMMAND} -E create_symlink "${ENGINE_SO_PATH}" "${_engine_symlink}" - ) -endif() - -# --------------------------------------------------------------------------- -# 3. Build the EP Python extension. -# --------------------------------------------------------------------------- -if("${EP_TORCH_VERSIONS}" STREQUAL "") - message(STATUS "[EP] Building with currently-installed PyTorch") - execute_process( - COMMAND ${Python3_EXECUTABLE} setup.py build_ext --build-lib . - WORKING_DIRECTORY "${SOURCE_DIR}" - RESULT_VARIABLE _ret - ) - if(NOT _ret EQUAL 0) - message(FATAL_ERROR "[EP] Extension build failed (exit code: ${_ret})") - endif() -else() - message(STATUS "[EP] Building for PyTorch versions: ${EP_TORCH_VERSIONS}") - foreach(_version IN LISTS EP_TORCH_VERSIONS) - install_pytorch_wheel("${_version}" "${EP_CUDA_MAJOR}" "${EP_CUDA_MINOR}" "[EP]") - - execute_process( - COMMAND ${Python3_EXECUTABLE} setup.py build_ext --build-lib . --force - WORKING_DIRECTORY "${SOURCE_DIR}" - RESULT_VARIABLE _ret - ) - if(NOT _ret EQUAL 0) - message(FATAL_ERROR "[EP] Extension build failed for PyTorch ${_version}") - endif() - endforeach() -endif() - -# --------------------------------------------------------------------------- -# 4. Copy the built .so files to the staging directory. -# --------------------------------------------------------------------------- -file(MAKE_DIRECTORY "${STAGING_DIR}") -file(GLOB _so_files "${SOURCE_DIR}/mooncake/*.so") -foreach(_so IN LISTS _so_files) - get_filename_component(_fname "${_so}" NAME) - message(STATUS "[EP] Staging ${_fname} -> ${STAGING_DIR}") - file(COPY "${_so}" DESTINATION "${STAGING_DIR}" NO_SOURCE_PERMISSIONS) -endforeach() - -message(STATUS "[EP] Mooncake EP extension build complete") diff --git a/mooncake-ep/CMakeLists.txt b/mooncake-ep/CMakeLists.txt index 11cf38c1f7..573748c127 100644 --- a/mooncake-ep/CMakeLists.txt +++ b/mooncake-ep/CMakeLists.txt @@ -1,33 +1,15 @@ cmake_minimum_required(VERSION 3.16) project(mooncake-ep) -# Find PyTorch's CMake prefix path -execute_process( - COMMAND ${PYTHON_EXECUTABLE} -c "import torch; print(torch.utils.cmake_prefix_path)" - OUTPUT_VARIABLE PYTORCH_CMAKE_PATH - OUTPUT_STRIP_TRAILING_WHITESPACE -) -if(NOT PYTORCH_CMAKE_PATH) - message(WARNING "Could not find PyTorch CMake path! Please set Torch_DIR.") -else () - message(STATUS "Found PyTorch CMake path: ${PYTORCH_CMAKE_PATH}") - list(APPEND CMAKE_PREFIX_PATH "${PYTORCH_CMAKE_PATH}/Torch") +find_package( + Python3 + COMPONENTS Interpreter Development.Module + REQUIRED) + +if(USE_CUDA) + enable_language(CUDA) + find_package(CUDAToolkit REQUIRED) endif() -set(TORCH_CUDA_ARCH_LIST "8.0;9.0") - -find_package(CUDAToolkit REQUIRED) -# https://discuss.pytorch.org/t/failed-to-find-nvtoolsext/179635/13 -if(NOT TARGET CUDA::nvToolsExt AND TARGET CUDA::nvtx3) - add_library(CUDA::nvToolsExt INTERFACE IMPORTED) - target_compile_definitions( - CUDA::nvToolsExt INTERFACE - TORCH_CUDA_USE_NVTX3 - ) - target_link_libraries(CUDA::nvToolsExt INTERFACE CUDA::nvtx3) -endif() -find_package(Torch REQUIRED) -include_directories(${TORCH_INCLUDE_DIRS}) - include_directories(include) add_subdirectory(src) diff --git a/mooncake-ep/benchmarks/legacy_buffer_perf.cpp b/mooncake-ep/benchmarks/legacy_buffer_perf.cpp new file mode 100644 index 0000000000..aa7b5776d3 --- /dev/null +++ b/mooncake-ep/benchmarks/legacy_buffer_perf.cpp @@ -0,0 +1,199 @@ +#include + +#include +#include +#include +#include +#include + +#include +#include +#include + +namespace py = pybind11; + +namespace mooncake { +namespace { + +struct LegacyBufferPerfTensors { + uint64_t x_ptr = 0; + uint64_t topk_idx_ptr = 0; + uint64_t topk_weights_ptr = 0; + uint64_t active_ranks_ptr = 0; + uint64_t expert_x_ptr = 0; + uint64_t packed_recv_x_ptr = 0; + uint64_t packed_recv_count_ptr = 0; + uint64_t packed_recv_src_info_ptr = 0; + uint64_t packed_recv_layout_range_ptr = 0; + uint64_t combined_x_ptr = 0; +}; + +struct LegacyBufferPerfConfig { + int num_tokens = 0; + int hidden = 0; + int num_topk = 0; + int num_local_experts = 0; + int num_max_dispatch_tokens_per_rank = 0; + int num_experts = 0; + int timeout_us = -1; + int warmups = 20; + int iterations = 30; + uint64_t compute_stream_ptr = 0; +}; + +struct LegacyBufferPerfResult { + double average_us = 0.0; + double min_us = 0.0; + double max_us = 0.0; +}; + +void cuda_check(cudaError_t status, const char* operation) { + if (status != cudaSuccess) { + throw std::runtime_error(std::string(operation) + ": " + + cudaGetErrorString(status)); + } +} + +void validate(MooncakeEpBuffer& buffer, const LegacyBufferPerfTensors& tensors, + const LegacyBufferPerfConfig& config) { + if (config.num_tokens <= 0 || config.hidden <= 0 || config.num_topk <= 0 || + config.num_max_dispatch_tokens_per_rank < config.num_tokens || + config.num_experts <= 0 || config.num_local_experts <= 0 || + config.warmups < 0 || config.iterations <= 1) { + throw std::invalid_argument( + "invalid legacy EP benchmark configuration"); + } + if (buffer.ibgda_disabled() && !buffer.use_fast_path()) { + throw std::runtime_error( + "legacy EP benchmark requires the native fast path"); + } + if (!tensors.x_ptr || !tensors.topk_idx_ptr || !tensors.topk_weights_ptr || + !tensors.active_ranks_ptr || !tensors.expert_x_ptr || + !tensors.packed_recv_x_ptr || !tensors.packed_recv_count_ptr || + !tensors.packed_recv_src_info_ptr || + !tensors.packed_recv_layout_range_ptr || !tensors.combined_x_ptr) { + throw std::invalid_argument( + "legacy EP benchmark received a null tensor pointer"); + } +} + +LegacyBufferPerfResult run_legacy_buffer_perf( + MooncakeEpBuffer& buffer, const LegacyBufferPerfTensors& tensors, + const LegacyBufferPerfConfig& config) { + validate(buffer, tensors, config); + + const int num_local_experts = config.num_local_experts; + const auto compute_stream = + reinterpret_cast(config.compute_stream_ptr); + const size_t recv_count_bytes = + static_cast(num_local_experts) * sizeof(int); + + auto run_once = [&] { + cuda_check(cudaMemsetAsync( + reinterpret_cast(tensors.packed_recv_count_ptr), + 0, recv_count_bytes, compute_stream), + "cudaMemsetAsync(packed_recv_count)"); + buffer.dispatch( + tensors.x_ptr, tensors.topk_idx_ptr, tensors.active_ranks_ptr, + config.num_tokens, config.hidden, config.num_topk, + config.num_max_dispatch_tokens_per_rank, config.num_experts, + config.timeout_us, false, tensors.packed_recv_x_ptr, 0, + tensors.packed_recv_count_ptr, tensors.packed_recv_src_info_ptr, + tensors.packed_recv_layout_range_ptr, false, false, + config.compute_stream_ptr); + buffer.combine( + tensors.expert_x_ptr, tensors.topk_idx_ptr, + tensors.topk_weights_ptr, tensors.packed_recv_src_info_ptr, + tensors.packed_recv_layout_range_ptr, tensors.active_ranks_ptr, + num_local_experts, config.num_tokens, config.hidden, + config.num_topk, config.num_max_dispatch_tokens_per_rank, + config.num_experts, config.timeout_us, false, + tensors.combined_x_ptr, false, false, config.compute_stream_ptr); + }; + + for (int i = 0; i < config.warmups; ++i) run_once(); + cuda_check(cudaStreamSynchronize(compute_stream), + "cudaStreamSynchronize(warmup)"); + + std::vector starts(config.iterations); + std::vector ends(config.iterations); + for (int i = 0; i < config.iterations; ++i) { + cuda_check(cudaEventCreate(&starts[i]), "cudaEventCreate(start)"); + cuda_check(cudaEventCreate(&ends[i]), "cudaEventCreate(end)"); + cuda_check(cudaEventRecord(starts[i], compute_stream), + "cudaEventRecord(start)"); + run_once(); + cuda_check(cudaEventRecord(ends[i], compute_stream), + "cudaEventRecord(end)"); + } + cuda_check(cudaEventSynchronize(ends.back()), "cudaEventSynchronize(end)"); + + std::vector timings_us; + timings_us.reserve(config.iterations - 1); + for (int i = 1; i < config.iterations; ++i) { + float elapsed_ms = 0.0f; + cuda_check(cudaEventElapsedTime(&elapsed_ms, starts[i], ends[i]), + "cudaEventElapsedTime"); + timings_us.push_back(elapsed_ms * 1000.0f); + } + for (auto event : starts) cudaEventDestroy(event); + for (auto event : ends) cudaEventDestroy(event); + + double total_us = 0.0; + for (float value : timings_us) total_us += value; + const auto [min_it, max_it] = + std::minmax_element(timings_us.begin(), timings_us.end()); + return {total_us / timings_us.size(), *min_it, *max_it}; +} + +} // namespace + +void bind_legacy_buffer_perf(py::module_& module) { + py::class_(module, "_LegacyBufferPerfTensors") + .def(py::init<>()) + .def_readwrite("x_ptr", &LegacyBufferPerfTensors::x_ptr) + .def_readwrite("topk_idx_ptr", &LegacyBufferPerfTensors::topk_idx_ptr) + .def_readwrite("topk_weights_ptr", + &LegacyBufferPerfTensors::topk_weights_ptr) + .def_readwrite("active_ranks_ptr", + &LegacyBufferPerfTensors::active_ranks_ptr) + .def_readwrite("expert_x_ptr", &LegacyBufferPerfTensors::expert_x_ptr) + .def_readwrite("packed_recv_x_ptr", + &LegacyBufferPerfTensors::packed_recv_x_ptr) + .def_readwrite("packed_recv_count_ptr", + &LegacyBufferPerfTensors::packed_recv_count_ptr) + .def_readwrite("packed_recv_src_info_ptr", + &LegacyBufferPerfTensors::packed_recv_src_info_ptr) + .def_readwrite("packed_recv_layout_range_ptr", + &LegacyBufferPerfTensors::packed_recv_layout_range_ptr) + .def_readwrite("combined_x_ptr", + &LegacyBufferPerfTensors::combined_x_ptr); + + py::class_(module, "_LegacyBufferPerfConfig") + .def(py::init<>()) + .def_readwrite("num_tokens", &LegacyBufferPerfConfig::num_tokens) + .def_readwrite("hidden", &LegacyBufferPerfConfig::hidden) + .def_readwrite("num_topk", &LegacyBufferPerfConfig::num_topk) + .def_readwrite("num_local_experts", + &LegacyBufferPerfConfig::num_local_experts) + .def_readwrite( + "num_max_dispatch_tokens_per_rank", + &LegacyBufferPerfConfig::num_max_dispatch_tokens_per_rank) + .def_readwrite("num_experts", &LegacyBufferPerfConfig::num_experts) + .def_readwrite("timeout_us", &LegacyBufferPerfConfig::timeout_us) + .def_readwrite("warmups", &LegacyBufferPerfConfig::warmups) + .def_readwrite("iterations", &LegacyBufferPerfConfig::iterations) + .def_readwrite("compute_stream_ptr", + &LegacyBufferPerfConfig::compute_stream_ptr); + + py::class_(module, "_LegacyBufferPerfResult") + .def_readonly("average_us", &LegacyBufferPerfResult::average_us) + .def_readonly("min_us", &LegacyBufferPerfResult::min_us) + .def_readonly("max_us", &LegacyBufferPerfResult::max_us); + + module.def("_benchmark_legacy_buffer", &run_legacy_buffer_perf, + py::arg("buffer"), py::arg("tensors"), py::arg("config"), + py::call_guard()); +} + +} // namespace mooncake diff --git a/mooncake-ep/benchmarks/legacy_buffer_perf.py b/mooncake-ep/benchmarks/legacy_buffer_perf.py new file mode 100644 index 0000000000..3a8aac09e4 --- /dev/null +++ b/mooncake-ep/benchmarks/legacy_buffer_perf.py @@ -0,0 +1,123 @@ +#!/usr/bin/env python3 +"""Native-core legacy EP benchmark with Python distributed bootstrap. + +Launch with torchrun. Python initializes the Mooncake process group and owns +the fixed tensor storage; the timed dispatch/combine loop executes entirely in +the torch-free C++ EP module. +""" + +import argparse +import os + +import torch +import torch.distributed as dist + +import mooncake._ep as native_ep +import mooncake.pg # Registers the Mooncake process-group backend for bootstrap. +from mooncake.mooncake_ep_buffer import Buffer + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--tokens", type=int, default=128) + parser.add_argument("--hidden", type=int, default=7168) + parser.add_argument("--experts", type=int, default=288) + parser.add_argument("--topk", type=int, default=8) + parser.add_argument("--warmups", type=int, default=20) + parser.add_argument("--iterations", type=int, default=30) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + dist.init_process_group("mooncake") + world_size = dist.get_world_size() + rank = dist.get_rank() + group = dist.new_group(list(range(world_size))) + + if args.experts % world_size: + raise ValueError("--experts must be divisible by world size") + if args.hidden % 128: + raise ValueError("--hidden must be divisible by 128") + + torch.manual_seed(rank) + num_local_experts = args.experts // world_size + buffer_bytes = Buffer.get_ep_buffer_size_hint( + args.tokens, args.hidden, world_size, args.experts + ) + buffer = Buffer(group, num_ep_buffer_bytes=buffer_bytes) + if buffer._use_fallback: + raise RuntimeError("native-core benchmark requires the EP fast path") + + x = torch.randn( + (args.tokens, args.hidden), dtype=torch.bfloat16, device="cuda" + ) + scores = torch.randn( + (args.tokens, args.experts), dtype=torch.float32, device="cuda" + ) + topk_idx = torch.topk(scores, args.topk, dim=-1).indices.contiguous() + topk_weights = torch.rand( + (args.tokens, args.topk), dtype=torch.float32, device="cuda" + ) + active_ranks = torch.ones(world_size, dtype=torch.int32, device="cuda") + recv_tokens = world_size * args.tokens + expert_x = torch.randn( + (num_local_experts, recv_tokens, args.hidden), + dtype=torch.bfloat16, + device="cuda", + ) + packed_recv_x = torch.empty_like(expert_x) + packed_recv_count = torch.empty( + num_local_experts, dtype=torch.int32, device="cuda" + ) + packed_recv_src_info = torch.empty( + (num_local_experts, recv_tokens), dtype=torch.int32, device="cuda" + ) + packed_recv_layout_range = torch.empty( + (num_local_experts, world_size), dtype=torch.int64, device="cuda" + ) + combined_x = torch.empty_like(x) + cache_flush = torch.empty(int(256e6 // 4), dtype=torch.int32, device="cuda") + + tensors = native_ep._LegacyBufferPerfTensors() + tensors.x_ptr = x.data_ptr() + tensors.topk_idx_ptr = topk_idx.data_ptr() + tensors.topk_weights_ptr = topk_weights.data_ptr() + tensors.active_ranks_ptr = active_ranks.data_ptr() + tensors.expert_x_ptr = expert_x.data_ptr() + tensors.packed_recv_x_ptr = packed_recv_x.data_ptr() + tensors.packed_recv_count_ptr = packed_recv_count.data_ptr() + tensors.packed_recv_src_info_ptr = packed_recv_src_info.data_ptr() + tensors.packed_recv_layout_range_ptr = packed_recv_layout_range.data_ptr() + tensors.combined_x_ptr = combined_x.data_ptr() + + config = native_ep._LegacyBufferPerfConfig() + config.num_tokens = args.tokens + config.hidden = args.hidden + config.num_topk = args.topk + config.num_local_experts = num_local_experts + config.num_max_dispatch_tokens_per_rank = args.tokens + config.num_experts = args.experts + config.warmups = args.warmups + config.iterations = args.iterations + config.compute_stream_ptr = torch.cuda.current_stream().cuda_stream + + dist.barrier(group=group) + cache_flush.zero_() + result = native_ep._benchmark_legacy_buffer(buffer.runtime, tensors, config) + selections = topk_idx.numel() + payload_bytes = selections * (args.hidden * 4) + bandwidth = payload_bytes / result.average_us / 1e3 + print( + f"[rank {rank}] Native dispatch + combine: {bandwidth:.2f} GB/s, " + f"avg_t={result.average_us:.2f} us, min_t={result.min_us:.2f} us, " + f"max_t={result.max_us:.2f} us", + flush=True, + ) + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/mooncake-ep/include/elastic/mooncake_ep_elastic_buffer.h b/mooncake-ep/include/elastic/mooncake_ep_elastic_buffer.h index 6b5dacdd07..4986c755fd 100644 --- a/mooncake-ep/include/elastic/mooncake_ep_elastic_buffer.h +++ b/mooncake-ep/include/elastic/mooncake_ep_elastic_buffer.h @@ -41,38 +41,6 @@ struct ElasticConfig { int num_gpu_timeout_secs = 100; }; -struct ElasticNativeHandle { - bool do_expand = false; - int num_experts = 0; - int expert_alignment = 1; - int num_max_tokens_per_rank = 0; - int num_sms = 0; - torch::Tensor topk_idx; - torch::Tensor psum_num_recv_tokens_per_scaleup_rank; - torch::Tensor psum_num_recv_tokens_per_expert; - torch::Tensor recv_src_metadata; - torch::Tensor recv_layout_range; - torch::Tensor dst_buffer_slot_idx; - std::optional token_metadata_at_forward; - std::optional channel_linked_list; - std::vector num_recv_tokens_per_expert_list; -}; - -struct ElasticDispatchOutput { - torch::Tensor recv_x; - std::optional recv_x_scales; - std::optional recv_topk_idx; - std::optional recv_topk_weights; - ElasticNativeHandle handle; - std::optional event; -}; - -struct ElasticCombineOutput { - torch::Tensor combined_x; - std::optional combined_topk_weights; - std::optional event; -}; - class MooncakeElasticBuffer { public: MooncakeElasticBuffer(int rank, int num_ranks, int64_t num_buffer_bytes, @@ -97,21 +65,30 @@ class MooncakeElasticBuffer { std::tuple get_logical_domain_size() const; int get_theoretical_num_sms(int num_experts, int num_topk) const; - ElasticDispatchOutput dispatch( - const torch::Tensor& x, const std::optional& sf, - const torch::Tensor& topk_idx, - const std::optional& topk_weights, - torch::Tensor& active_ranks, int num_experts, - int num_max_tokens_per_rank, int expert_alignment, int num_sms, - bool do_expand, bool do_cpu_sync, bool async_with_compute_stream, - const std::optional& cached_handle = std::nullopt); - - ElasticCombineOutput combine( - const torch::Tensor& x, const ElasticNativeHandle& handle, - const std::optional& topk_weights, - torch::Tensor& active_ranks, int num_sms, - bool async_with_compute_stream, - const std::optional& out); + std::optional dispatch( + uint64_t x_ptr, int x_element_size, uint64_t sf_ptr, int num_tokens, + int hidden, int num_sf_packs, int sf_token_stride, int sf_hidden_stride, + uint64_t topk_idx_ptr, int num_topk, uint64_t topk_weights_ptr, + uint64_t active_ranks_ptr, int num_experts, int num_max_tokens_per_rank, + int expert_alignment, int num_sms, bool do_expand, + bool async_with_compute_stream, uint64_t compute_stream_ptr, + bool cached_mode, uint64_t psum_num_recv_tokens_per_scaleup_rank_ptr, + uint64_t psum_num_recv_tokens_per_expert_ptr, + uint64_t dst_buffer_slot_idx_ptr, + uint64_t token_metadata_at_forward_ptr, + uint64_t channel_linked_list_ptr, uint64_t recv_x_ptr, + uint64_t recv_x_scales_ptr, uint64_t recv_topk_idx_ptr, + uint64_t recv_topk_weights_ptr, uint64_t recv_src_metadata_ptr); + + std::optional combine( + uint64_t x_ptr, int num_input_tokens, int hidden, uint64_t topk_idx_ptr, + int num_combined_tokens, int num_topk, uint64_t topk_weights_ptr, + uint64_t psum_num_recv_tokens_per_scaleup_rank_ptr, + uint64_t recv_src_metadata_ptr, uint64_t token_metadata_at_forward_ptr, + uint64_t channel_linked_list_ptr, uint64_t active_ranks_ptr, + int num_experts, int num_max_tokens_per_rank, bool do_expand, + int num_sms, bool async_with_compute_stream, + uint64_t compute_stream_ptr, uint64_t combined_x_ptr); MooncakeEpBuffer& native_buffer() { return *native_buffer_; } @@ -157,12 +134,15 @@ class MooncakeElasticBuffer { int64_t host_workspace_bytes_ = 0; void* host_workspace_ = nullptr; void* mapped_host_workspace_ = nullptr; + std::shared_ptr deterministic_rank_count_buffer_; + int64_t deterministic_rank_count_buffer_bytes_ = 0; static ElasticLaunchContext make_launch_context( MooncakeEpBuffer& buffer, const ElasticTopology& topology, void* mapped_host_workspace, int64_t timeout_cycles); static ElasticTopology discover_topology(int rank, int num_ranks, bool allow_hybrid_mode); + std::shared_ptr ensure_deterministic_rank_count_buffer(int num_sms); }; } // namespace mooncake diff --git a/mooncake-ep/include/elastic/mooncake_ep_elastic_compiled.cuh b/mooncake-ep/include/elastic/mooncake_ep_elastic_compiled.cuh index 1850361106..e5b277f8f6 100644 --- a/mooncake-ep/include/elastic/mooncake_ep_elastic_compiled.cuh +++ b/mooncake-ep/include/elastic/mooncake_ep_elastic_compiled.cuh @@ -28,8 +28,7 @@ #endif #include -#include -#include +#include #if defined(MOONCAKE_EP_USE_MUSA) && defined(__MCC__) && \ !defined(MOONCAKE_EP_MUSA_LDG_DEFINED) @@ -40,7 +39,7 @@ __device__ __forceinline__ dtype_t __ldg(const dtype_t* ptr) { } #endif -#ifndef DISABLE_SM90_FEATURES +#if !defined(MOONCAKE_EP_USE_MUSA) && !defined(DISABLE_SM90_FEATURES) #include #elif !defined(MOONCAKE_EP_USE_MUSA) // Ampere does not support FP8 features diff --git a/mooncake-ep/include/elastic/mooncake_ep_elastic_exception.cuh b/mooncake-ep/include/elastic/mooncake_ep_elastic_exception.cuh index 26e1c8b984..371f16ee8f 100644 --- a/mooncake-ep/include/elastic/mooncake_ep_elastic_exception.cuh +++ b/mooncake-ep/include/elastic/mooncake_ep_elastic_exception.cuh @@ -39,7 +39,7 @@ #endif #ifndef EP_UNIFIED_ASSERT -#ifdef __CUDA_ARCH__ +#if defined(__CUDA_ARCH__) || defined(__MUSA_ARCH__) #define EP_UNIFIED_ASSERT(cond) EP_DEVICE_ASSERT(cond) #else #define EP_UNIFIED_ASSERT(cond) EP_HOST_ASSERT(cond) diff --git a/mooncake-ep/include/elastic/mooncake_ep_elastic_launch.cuh b/mooncake-ep/include/elastic/mooncake_ep_elastic_launch.cuh index e21c86a170..e274373fec 100644 --- a/mooncake-ep/include/elastic/mooncake_ep_elastic_launch.cuh +++ b/mooncake-ep/include/elastic/mooncake_ep_elastic_launch.cuh @@ -2,8 +2,7 @@ #include -#include -#include +#include namespace mooncake { diff --git a/mooncake-ep/include/elastic/mooncake_ep_elastic_ptx.cuh b/mooncake-ep/include/elastic/mooncake_ep_elastic_ptx.cuh index f1fa20650e..2047b2e117 100644 --- a/mooncake-ep/include/elastic/mooncake_ep_elastic_ptx.cuh +++ b/mooncake-ep/include/elastic/mooncake_ep_elastic_ptx.cuh @@ -3,7 +3,7 @@ // transport references are replaced with Mooncake Device API adapters. #pragma once -#include +#include #include #include @@ -22,7 +22,7 @@ using arrival_phase = uint32_t; // More than TMA, `longlong4` requires 32 bytes aligned static constexpr int kNumTMAAlignBytes = 32; -#ifdef __CUDACC__ +#if defined(__CUDACC__) || defined(__MUSACC__) /// Exceptions __forceinline__ __device__ void trap() { diff --git a/mooncake-ep/include/mooncake_ep_api.cuh b/mooncake-ep/include/mooncake_ep_api.cuh index 1a560d1b90..bc747c2f22 100644 --- a/mooncake-ep/include/mooncake_ep_api.cuh +++ b/mooncake-ep/include/mooncake_ep_api.cuh @@ -1,6 +1,6 @@ #pragma once -#include +#include namespace mooncake { diff --git a/mooncake-ep/include/mooncake_ep_buffer.h b/mooncake-ep/include/mooncake_ep_buffer.h index 9307936c80..7bcf5140de 100644 --- a/mooncake-ep/include/mooncake_ep_buffer.h +++ b/mooncake-ep/include/mooncake_ep_buffer.h @@ -1,16 +1,15 @@ #ifndef MOONCAKE_EP_BUFFER_H #define MOONCAKE_EP_BUFFER_H -#include -#include -#include -#include +#include +#include #include +#include #include #include #include #include -#include +#include #include namespace mooncake { @@ -72,7 +71,7 @@ struct MooncakeEpBuffer { int rank, num_ranks; int clock_rate_khz; - // GDR buffer — owned by p2p_transport_ + // GDR buffer — allocated by p2p_transport_; peer mappings are optional. int buffer_idx{}; int phase_epochs[2]{}; int64_t num_ep_buffer_bytes; @@ -90,16 +89,16 @@ struct MooncakeEpBuffer { std::unique_ptr owned_rdma_transport_; bool ibgda_disabled_ = false; + bool p2p_enabled_ = true; int USE_QP_COUNT = MAX_QP_COUNT; - // Cap on active RoCE QPs per peer: spreading small EP messages across too - // many QP/doorbell/progress streams hurts when GPUs share an HCA. Default - // 8; override at runtime with MOONCAKE_EP_ACTIVE_QPS_PER_RANK (>= per-rank - // QP count disables). - int active_qps_cap_ = 8; + // Active RoCE QPs per peer. The platform-specific default is selected in + // active_qps_per_rank_for_ep(); a positive + // MOONCAKE_EP_ACTIVE_QPS_PER_RANK value forces an explicit count. + int active_qps_cap_ = 0; // Stream for communication - at::cuda::CUDAStream comm_stream; + cudaStream_t comm_stream = nullptr; // Workspace void* workspace = nullptr; @@ -109,31 +108,32 @@ struct MooncakeEpBuffer { // (engine owns the transports). Otherwise EP creates its own via the // factory functions (EP owns them via owned_p2p_transport_ etc.). MooncakeEpBuffer(int rank, int num_ranks, int64_t num_ep_buffer_bytes, + bool disable_p2p = false, TransferEngine* engine = nullptr); ~MooncakeEpBuffer() noexcept(false); - std::tuple, torch::Tensor, - torch::Tensor, torch::Tensor, std::optional, - std::optional>> - dispatch(const torch::Tensor& x, const torch::Tensor& topk_idx, - torch::Tensor& active_ranks, int num_max_dispatch_tokens_per_rank, - int num_experts, int timeout_us, bool use_fp8, bool async, - bool return_recv_hook); - - std::tuple, - std::optional>> - combine(const torch::Tensor& x, const torch::Tensor& topk_idx, - const torch::Tensor& topk_weights, const torch::Tensor& src_info, - const torch::Tensor& layout_range, torch::Tensor& active_ranks, + std::tuple, std::optional>> + dispatch(uint64_t x_ptr, uint64_t topk_idx_ptr, uint64_t active_ranks_ptr, + int num_tokens, int hidden, int num_topk, + int num_max_dispatch_tokens_per_rank, int num_experts, + int timeout_us, bool use_fp8, uint64_t packed_recv_x_ptr, + uint64_t packed_recv_x_scales_ptr, uint64_t packed_recv_count_ptr, + uint64_t packed_recv_src_info_ptr, + uint64_t packed_recv_layout_range_ptr, bool async, + bool return_recv_hook, uint64_t compute_stream_ptr); + + std::tuple, std::optional>> + combine(uint64_t x_ptr, uint64_t topk_idx_ptr, uint64_t topk_weights_ptr, + uint64_t src_info_ptr, uint64_t layout_range_ptr, + uint64_t active_ranks_ptr, int num_local_experts, + int num_combined_tokens, int hidden, int num_topk, int num_max_dispatch_tokens_per_rank, int num_experts, - int timeout_us, bool zero_copy, bool async, bool return_recv_hook, - const std::optional& out); - - torch::Tensor get_next_combine_buffer(int num_max_dispatch_tokens_per_rank, - int hidden, int num_experts); + int timeout_us, bool zero_copy, uint64_t combined_x_ptr, bool async, + bool return_recv_hook, uint64_t compute_stream_ptr); bool ibgda_disabled() const { return ibgda_disabled_; } + bool p2p_enabled() const { return p2p_enabled_; } bool is_roce() const { return rdma_transport_ && rdma_transport_->isRoce(); diff --git a/mooncake-ep/include/mooncake_ep_configs.cuh b/mooncake-ep/include/mooncake_ep_configs.cuh index 1e7f0c2149..5d05908f36 100644 --- a/mooncake-ep/include/mooncake_ep_configs.cuh +++ b/mooncake-ep/include/mooncake_ep_configs.cuh @@ -39,17 +39,23 @@ #undef __CUDA_NO_BFLOAT162_OPERATORS__ #endif +#include +#if !defined(MOONCAKE_EP_USE_MUSA) && !defined(MOONCAKE_EP_USE_MACA) #include -#ifndef MOONCAKE_EP_USE_MACA #include +#endif +#ifndef MOONCAKE_EP_USE_MACA #include #endif -#include #if defined(MOONCAKE_EP_USE_MUSA) || defined(MOONCAKE_EP_USE_MACA) #define MOONCAKE_EP_SPLIT_SEND_RECV 1 #endif +#if defined(MOONCAKE_EP_USE_MACA) +#define MOONCAKE_EP_PHASE_ACK 1 +#endif + // torchada maps nv_bfloat16 → __mt_bfloat16 which is an incomplete type on // MUSA, so sizeof(__mt_bfloat16) fails. mt_bfloat16 (the complete typedef in // musa_bf16.hpp) requires the MUSA device compiler (mcc) and cannot be diff --git a/mooncake-ep/include/mooncake_ep_device.h b/mooncake-ep/include/mooncake_ep_device.h index e88417ff7f..0ab7b58f1d 100644 --- a/mooncake-ep/include/mooncake_ep_device.h +++ b/mooncake-ep/include/mooncake_ep_device.h @@ -13,14 +13,14 @@ #include using ep_fp8_storage_t = __mt_fp8_storage_t; using ep_fp8x2_storage_t = __mt_fp8x2_storage_t; -#if defined(__CUDACC__) || defined(__MCC__) +#if defined(__CUDACC__) || defined(__MCC__) || defined(__MUSACC__) __device__ __forceinline__ ep_fp8x2_storage_t ep_cvt_float2_to_fp8x2(float2 x) { return __musa_cvt_float2_to_fp8x2(x, __MT_SATFINITE, __MT_E4M3); } #endif // -- Device intrinsics (MUSA doesn't have __ldg / __activemask) -------------- -#if (defined(__CUDACC__) || defined(__MCC__)) && \ +#if (defined(__CUDACC__) || defined(__MCC__) || defined(__MUSACC__)) && \ !defined(MOONCAKE_EP_MUSA_LDG_DEFINED) #define MOONCAKE_EP_MUSA_LDG_DEFINED template @@ -32,7 +32,7 @@ __device__ __forceinline__ dtype_t __ldg(const dtype_t* ptr) { #define __activemask() (0xffffffff) #endif -#if defined(__CUDACC__) || defined(__MCC__) +#if defined(__CUDACC__) || defined(__MCC__) || defined(__MUSACC__) __forceinline__ __device__ int get_lane_id() { return threadIdx.x % 32; } #endif @@ -44,15 +44,11 @@ __forceinline__ __device__ int get_lane_id() { return threadIdx.x % 32; } dim3 _block(num_threads); \ cudaStream_t _stream = stream -#define LAUNCH_KERNEL(config, kernel, ...) \ - kernel<<<_grid, _block, 0, _stream>>>(__VA_ARGS__); \ - { \ - auto _err = cudaGetLastError(); \ - if (_err != cudaSuccess) { \ - fprintf(stderr, "[EP] kernel launch failed: %s\n", \ - cudaGetErrorString(_err)); \ - } \ - } +#define LAUNCH_KERNEL(config, kernel, ...) \ + do { \ + kernel<<<_grid, _block, 0, _stream>>>(__VA_ARGS__); \ + CUDA_CHECK(cudaGetLastError()); \ + } while (false) #elif defined(MOONCAKE_EP_USE_MACA) @@ -63,7 +59,7 @@ __forceinline__ __device__ int get_lane_id() { return threadIdx.x % 32; } #include using ep_fp8_storage_t = uint8_t; using ep_fp8x2_storage_t = uint16_t; -#if defined(__CUDACC__) || defined(__MCC__) +#if defined(__CUDACC__) || defined(__MCC__) || defined(__MUSACC__) __device__ __forceinline__ ep_fp8x2_storage_t ep_cvt_float2_to_fp8x2(float2) { return 0; } @@ -74,7 +70,7 @@ __device__ __forceinline__ ep_fp8x2_storage_t ep_cvt_float2_to_fp8x2(float2) { #define __activemask() (0xffffffff) #endif -#if defined(__CUDACC__) || defined(__MCC__) +#if defined(__CUDACC__) || defined(__MCC__) || defined(__MUSACC__) __forceinline__ __device__ int get_lane_id() { return threadIdx.x % 32; } #endif @@ -86,15 +82,11 @@ __forceinline__ __device__ int get_lane_id() { return threadIdx.x % 32; } dim3 _block(num_threads); \ cudaStream_t _stream = stream -#define LAUNCH_KERNEL(config, kernel, ...) \ - kernel<<<_grid, _block, 0, _stream>>>(__VA_ARGS__); \ - { \ - auto _err = cudaGetLastError(); \ - if (_err != cudaSuccess) { \ - fprintf(stderr, "[EP] kernel launch failed: %s\n", \ - cudaGetErrorString(_err)); \ - } \ - } +#define LAUNCH_KERNEL(config, kernel, ...) \ + do { \ + kernel<<<_grid, _block, 0, _stream>>>(__VA_ARGS__); \ + CUDA_CHECK(cudaGetLastError()); \ + } while (false) #else // !MOONCAKE_EP_USE_MUSA && !MOONCAKE_EP_USE_MACA diff --git a/mooncake-ep/include/mooncake_ep_event.h b/mooncake-ep/include/mooncake_ep_event.h index 4809704983..c0e8a055e4 100644 --- a/mooncake-ep/include/mooncake_ep_event.h +++ b/mooncake-ep/include/mooncake_ep_event.h @@ -1,49 +1,62 @@ #pragma once -#include +#include #include #include -#include namespace mooncake { struct EventHandle { - std::shared_ptr event; + std::shared_ptr event; + std::shared_ptr keepalive; EventHandle() { - event = std::make_shared(torch::kCUDA); - event->record(at::cuda::getCurrentCUDAStream()); + event = std::shared_ptr(new cudaEvent_t(nullptr), + [](cudaEvent_t* p) { + if (p != nullptr) { + if (*p != nullptr) + cudaEventDestroy(*p); + delete p; + } + }); + CUDA_CHECK( + cudaEventCreateWithFlags(event.get(), cudaEventDisableTiming)); } - explicit EventHandle(const at::cuda::CUDAStream& stream) { - event = std::make_shared(torch::kCUDA); - event->record(stream); + explicit EventHandle(uint64_t stream_ptr, + std::shared_ptr keepalive = nullptr) + : EventHandle() { + this->keepalive = std::move(keepalive); + auto stream = reinterpret_cast(stream_ptr); + CUDA_CHECK(cudaEventRecord(*event, stream)); } EventHandle(const EventHandle& other) = default; - void current_stream_wait() const { - at::cuda::getCurrentCUDAStream().unwrap().wait(*event); + void current_stream_wait(uint64_t stream_ptr) const { + auto stream = reinterpret_cast(stream_ptr); + CUDA_CHECK(cudaStreamWaitEvent(stream, *event, 0)); } - void synchronize() const { event->synchronize(); } + void synchronize() const { CUDA_CHECK(cudaEventSynchronize(*event)); } }; -inline torch::Event create_event(const at::cuda::CUDAStream& s) { - auto event = torch::Event(torch::kCUDA); - event.record(s); +inline cudaEvent_t create_event(cudaStream_t stream) { + cudaEvent_t event = nullptr; + CUDA_CHECK(cudaEventCreateWithFlags(&event, cudaEventDisableTiming)); + CUDA_CHECK(cudaEventRecord(event, stream)); return event; } -inline void stream_wait(const at::cuda::CUDAStream& s_0, - const at::cuda::CUDAStream& s_1) { - EP_HOST_ASSERT(s_0.id() != s_1.id()); - s_0.unwrap().wait(create_event(s_1)); +inline void stream_wait(cudaStream_t dst_stream, cudaStream_t src_stream) { + EP_HOST_ASSERT(dst_stream != src_stream); + auto event = create_event(src_stream); + CUDA_CHECK(cudaStreamWaitEvent(dst_stream, event, 0)); + CUDA_CHECK(cudaEventDestroy(event)); } -inline void stream_wait(const at::cuda::CUDAStream& s, - const EventHandle& event) { - s.unwrap().wait(*event.event); +inline void stream_wait(cudaStream_t s, const EventHandle& event) { + CUDA_CHECK(cudaStreamWaitEvent(s, *event.event, 0)); } } // namespace mooncake diff --git a/mooncake-ep/setup.py b/mooncake-ep/setup.py deleted file mode 100644 index 2a4f388e1c..0000000000 --- a/mooncake-ep/setup.py +++ /dev/null @@ -1,138 +0,0 @@ -import os -import re - -from setuptools import setup -import torch - -use_musa = os.getenv("MOONCAKE_EP_USE_MUSA", "").upper() in {"1", "ON", "TRUE", "YES"} -use_maca = ( - os.getenv("MOONCAKE_EP_USE_MACA", "").upper() in {"1", "ON", "TRUE", "YES"} - or (hasattr(torch.version, "maca") and torch.version.maca is not None) -) -if use_musa: - try: - import importlib - - importlib.import_module("torchada") - except ImportError as e: - raise ImportError( - "torchada is required to build the MUSA EP extension. " - "Please install it first using 'pip install torchada'." - ) from e - - -from torch.utils.cpp_extension import ( # noqa: E402 - BuildExtension, - CUDAExtension, - CUDA_HOME, -) - - -torch_version = re.match(r"\d+(?:\.\d+)*", torch.__version__).group() -version_suffix = "_" + torch_version.replace(".", "_") -module_name = "mooncake.ep" + version_suffix - -abi_flag = int(torch._C._GLIBCXX_USE_CXX11_ABI) -current_dir = os.path.abspath(os.path.dirname(__file__)) -repo_dir = os.path.abspath(os.path.join(current_dir, os.pardir)) -sysroot_dir = os.path.join(repo_dir, ".deps", "sysroot", "usr") - - -def existing_dirs(*paths): - return [path for path in paths if os.path.isdir(path)] - - -sysroot_include_dirs = existing_dirs( - os.path.join(sysroot_dir, "include"), - os.path.join(sysroot_dir, "include", "jsoncpp"), - os.path.join(sysroot_dir, "include", "libnl3"), -) -sysroot_library_dirs = existing_dirs( - os.path.join(sysroot_dir, "lib", "x86_64-linux-gnu"), - os.path.join(sysroot_dir, "lib"), -) - -abi_define = f"-D_GLIBCXX_USE_CXX11_ABI={abi_flag}" -cxx_args = [abi_define, "-std=c++20", "-O3", "-g0"] - -cuda_libraries = ["ibverbs", "mlx5"] -cuda_library_dirs = [] -include_dirs = [ - os.path.join(current_dir, "include"), - os.path.join(current_dir, "../mooncake-transfer-engine/include"), -] - -if use_musa: - cuda_libraries = [] - musa_defines = [ - "-DUSE_MUSA", - "-DMOONCAKE_EP_USE_MUSA=1", - ] - cxx_args += musa_defines - # torchada maps the "nvcc" key to "mcc". - device_args = [ - abi_define, - *musa_defines, - "-std=c++20", - "--cuda-gpu-arch=mp_21", - "--cuda-gpu-arch=mp_31", - "-O3", - ] -elif use_maca: - cuda_libraries = [] - cuda_library_dirs = sysroot_library_dirs.copy() - include_dirs += sysroot_include_dirs - maca_defines = ["-DUSE_MACA", "-DMOONCAKE_EP_USE_MACA=1"] - cxx_args += maca_defines - device_args = [ - abi_define, - *maca_defines, - "-std=c++20", - "-O3", - ] -else: - cxx_args.append("-DUSE_CUDA") - device_args = [ - abi_define, - "-std=c++20", - "-DUSE_CUDA", - "-Xcompiler", - "-O3", - "-Xcompiler", - "-g0", - ] - # Link against the CUDA driver stub library if available. - if CUDA_HOME is not None: - cuda_stub_dir = os.path.join(CUDA_HOME, "lib64", "stubs") - cuda_stub_lib = os.path.join(cuda_stub_dir, "libcuda.so") - if os.path.exists(cuda_stub_lib): - cuda_libraries.insert(0, "cuda") - cuda_library_dirs.append(cuda_stub_dir) - -setup( - name=module_name, - ext_modules=[ - CUDAExtension( - name=module_name, - include_dirs=include_dirs, - sources=[ - "src/ep_py.cpp", - "src/mooncake_ep_buffer.cpp", - "src/mooncake_ep_elastic_buffer.cpp", - "src/mooncake_ep_kernel.cu", - "src/mooncake_ep_elastic_kernel.cu", - ], - extra_compile_args={"cxx": cxx_args, "nvcc": device_args}, - libraries=cuda_libraries, - library_dirs=cuda_library_dirs, - extra_link_args=[ - "-Wl,-rpath,$ORIGIN", - "-L" + os.path.join(current_dir, "../mooncake-wheel/mooncake"), - "-Wl,--push-state,--no-as-needed", - "-l:engine.so", - "-Wl,--pop-state", - ], - ) - ], - cmdclass={"build_ext": BuildExtension}, -) diff --git a/mooncake-ep/src/CMakeLists.txt b/mooncake-ep/src/CMakeLists.txt index a102f5011e..57ec0e2ed6 100644 --- a/mooncake-ep/src/CMakeLists.txt +++ b/mooncake-ep/src/CMakeLists.txt @@ -1,4 +1,95 @@ -add_library(mooncake_ep ep_py.cpp mooncake_ep_buffer.cpp mooncake_ep_elastic_buffer.cpp mooncake_ep_kernel.cu mooncake_ep_elastic_kernel.cu) +set(MOONCAKE_EP_SOURCES + ep_py.cpp + ../benchmarks/legacy_buffer_perf.cpp + mooncake_ep_buffer.cpp + mooncake_ep_elastic_buffer.cpp) -set_target_properties(mooncake_ep PROPERTIES POSITION_INDEPENDENT_CODE ON) -target_link_libraries(mooncake_ep PUBLIC ${TORCH_LIBRARIES} transfer_engine ibverbs mlx5) +set(MOONCAKE_EP_DEVICE_SOURCES + "${CMAKE_CURRENT_SOURCE_DIR}/mooncake_ep_kernel.cu" + "${CMAKE_CURRENT_SOURCE_DIR}/mooncake_ep_elastic_kernel.cu") + +if(USE_CUDA) + enable_language(CUDA) + list(APPEND MOONCAKE_EP_SOURCES ${MOONCAKE_EP_DEVICE_SOURCES}) +elseif(USE_MUSA) + if(DEFINED ENV{MUSA_HOME} AND NOT "$ENV{MUSA_HOME}" STREQUAL "") + set(_ep_musa_compiler_hint "$ENV{MUSA_HOME}/bin") + else() + set(_ep_musa_compiler_hint /usr/local/musa/bin) + endif() + find_program(_ep_musa_compiler NAMES mcc HINTS "${_ep_musa_compiler_hint}") + if(NOT _ep_musa_compiler) + message(FATAL_ERROR "USE_MUSA=ON requires the MUSA compiler (mcc)") + endif() + set(_ep_musa_depfile_supported FALSE) + if(CMAKE_GENERATOR MATCHES "^Ninja" OR + (CMAKE_GENERATOR MATCHES "Makefiles" AND + CMAKE_VERSION VERSION_GREATER_EQUAL 3.20)) + set(_ep_musa_depfile_supported TRUE) + endif() + + foreach(_ep_device_source IN LISTS MOONCAKE_EP_DEVICE_SOURCES) + get_filename_component(_ep_device_name "${_ep_device_source}" NAME_WE) + set(_ep_device_object + "${CMAKE_CURRENT_BINARY_DIR}/${_ep_device_name}_musa.o") + set(_ep_device_depfile + "${CMAKE_CURRENT_BINARY_DIR}/${_ep_device_name}_musa.d") + set(_ep_musa_depfile_argument) + if(_ep_musa_depfile_supported) + set(_ep_musa_depfile_argument DEPFILE "${_ep_device_depfile}") + endif() + add_custom_command( + OUTPUT "${_ep_device_object}" + COMMAND "${_ep_musa_compiler}" + -x musa + -std=c++20 + -O3 + -fPIC + -DUSE_MUSA + -DMOONCAKE_EP_USE_MUSA=1 + -MMD + -MT "${_ep_device_object}" + -MF "${_ep_device_depfile}" + --cuda-gpu-arch=mp_21 + --cuda-gpu-arch=mp_31 + "-I${CMAKE_CURRENT_SOURCE_DIR}/../include" + "-I${CMAKE_CURRENT_SOURCE_DIR}/../../mooncake-transfer-engine/include" + -c "${_ep_device_source}" -o "${_ep_device_object}" + DEPENDS "${_ep_device_source}" + ${_ep_musa_depfile_argument} + COMMENT "Compiling Mooncake EP MUSA device source ${_ep_device_name}" + VERBATIM) + set_source_files_properties("${_ep_device_object}" + PROPERTIES GENERATED TRUE EXTERNAL_OBJECT TRUE) + list(APPEND MOONCAKE_EP_SOURCES "${_ep_device_object}") + endforeach() +endif() + +pybind11_add_module(_ep MODULE ${MOONCAKE_EP_SOURCES}) +set_target_properties(_ep PROPERTIES POSITION_INDEPENDENT_CODE ON) +set_target_properties(_ep PROPERTIES INSTALL_RPATH "$ORIGIN") +if(USE_CUDA) + set_target_properties(_ep PROPERTIES CUDA_ARCHITECTURES "80;90") +endif() +if(USE_MUSA) + target_compile_definitions(_ep PRIVATE MOONCAKE_EP_USE_MUSA=1) +endif() +if(USE_MACA) + target_compile_definitions(_ep PRIVATE MOONCAKE_EP_USE_MACA=1) +endif() + +target_include_directories(_ep PRIVATE ${Python3_INCLUDE_DIRS}) +target_link_libraries(_ep PRIVATE transfer_engine ibverbs mlx5 glog::glog + gflags::gflags) + +if(USE_CUDA) + target_compile_options( + _ep PRIVATE + $<$:-Xcompiler=-O3> + $<$:-Xcompiler=-g0> + $<$:--expt-relaxed-constexpr>) + target_link_libraries(_ep PRIVATE CUDA::cudart) +elseif(USE_MUSA) + set_target_properties(_ep PROPERTIES LINKER_LANGUAGE CXX) + target_link_libraries(_ep PRIVATE musa musart rt) +endif() diff --git a/mooncake-ep/src/ep_py.cpp b/mooncake-ep/src/ep_py.cpp index 02307c0caf..a6ebe53932 100644 --- a/mooncake-ep/src/ep_py.cpp +++ b/mooncake-ep/src/ep_py.cpp @@ -4,72 +4,33 @@ #include #include #include -#include -#include -#include namespace py = pybind11; namespace mooncake { -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { +void bind_legacy_buffer_perf(py::module_& module); + +PYBIND11_MODULE(_ep, m) { m.def("get_ep_buffer_size_hint", &get_ep_buffer_size_hint); m.def("calculate_elastic_buffer_size", &MooncakeElasticBuffer::calculate_buffer_size); py::class_(m, "EventHandle") - .def(py::init<>()) - .def("current_stream_wait", &EventHandle::current_stream_wait) + .def(py::init(), py::arg("stream_ptr") = 0) + .def("current_stream_wait", &EventHandle::current_stream_wait, + py::arg("stream_ptr")) .def("synchronize", &EventHandle::synchronize); - py::class_(m, "ElasticNativeHandle") - .def(py::init<>()) - .def_readwrite("do_expand", &ElasticNativeHandle::do_expand) - .def_readwrite("num_experts", &ElasticNativeHandle::num_experts) - .def_readwrite("expert_alignment", - &ElasticNativeHandle::expert_alignment) - .def_readwrite("num_max_tokens_per_rank", - &ElasticNativeHandle::num_max_tokens_per_rank) - .def_readwrite("num_sms", &ElasticNativeHandle::num_sms) - .def_readwrite("topk_idx", &ElasticNativeHandle::topk_idx) - .def_readwrite( - "psum_num_recv_tokens_per_scaleup_rank", - &ElasticNativeHandle::psum_num_recv_tokens_per_scaleup_rank) - .def_readwrite("psum_num_recv_tokens_per_expert", - &ElasticNativeHandle::psum_num_recv_tokens_per_expert) - .def_readwrite("recv_src_metadata", - &ElasticNativeHandle::recv_src_metadata) - .def_readwrite("recv_layout_range", - &ElasticNativeHandle::recv_layout_range) - .def_readwrite("dst_buffer_slot_idx", - &ElasticNativeHandle::dst_buffer_slot_idx) - .def_readwrite("token_metadata_at_forward", - &ElasticNativeHandle::token_metadata_at_forward) - .def_readwrite("channel_linked_list", - &ElasticNativeHandle::channel_linked_list) - .def_readwrite("num_recv_tokens_per_expert_list", - &ElasticNativeHandle::num_recv_tokens_per_expert_list); - - py::class_(m, "ElasticDispatchOutput") - .def_readonly("recv_x", &ElasticDispatchOutput::recv_x) - .def_readonly("recv_x_scales", &ElasticDispatchOutput::recv_x_scales) - .def_readonly("recv_topk_idx", &ElasticDispatchOutput::recv_topk_idx) - .def_readonly("recv_topk_weights", - &ElasticDispatchOutput::recv_topk_weights) - .def_readonly("handle", &ElasticDispatchOutput::handle) - .def_readonly("event", &ElasticDispatchOutput::event); - - py::class_(m, "ElasticCombineOutput") - .def_readonly("combined_x", &ElasticCombineOutput::combined_x) - .def_readonly("combined_topk_weights", - &ElasticCombineOutput::combined_topk_weights) - .def_readonly("event", &ElasticCombineOutput::event); - m.attr("MAX_QP_COUNT") = pybind11::int_(MAX_QP_COUNT); + bind_legacy_buffer_perf(m); py::class_(m, "Buffer") - .def(py::init()) + .def(py::init(), py::arg("rank"), + py::arg("num_ranks"), py::arg("num_ep_buffer_bytes"), + py::arg("disable_p2p") = false) .def("ibgda_disabled", &MooncakeEpBuffer::ibgda_disabled) + .def("p2p_enabled", &MooncakeEpBuffer::p2p_enabled) .def("use_fast_path", &MooncakeEpBuffer::use_fast_path) .def("update_local_qpns", &MooncakeEpBuffer::update_local_qpns) .def("is_roce", &MooncakeEpBuffer::is_roce) @@ -82,9 +43,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { .def("sync_nvlink_ipc_handles", &MooncakeEpBuffer::sync_nvlink_ipc_handles) .def("dispatch", &MooncakeEpBuffer::dispatch) - .def("combine", &MooncakeEpBuffer::combine) - .def("get_next_combine_buffer", - &MooncakeEpBuffer::get_next_combine_buffer); + .def("combine", &MooncakeEpBuffer::combine); py::class_(m, "ElasticBuffer") .def(py::init= the per-rank QP count - // (e.g. 256) to effectively disable the cap. + // Optional runtime override for the RoCE active-QP count. Without an + // override, CUDA uses eight QPs and MUSA scales up to local experts. + active_qps_cap_ = 0; if (const char* env = std::getenv("MOONCAKE_EP_ACTIVE_QPS_PER_RANK")) { char* end = nullptr; long v = std::strtol(env, &end, 10); @@ -65,14 +95,17 @@ MooncakeEpBuffer::MooncakeEpBuffer(int rank, int num_ranks, << env << "'"; } } - LOG(INFO) << "[EP] RoCE active QPs/rank cap = " << active_qps_cap_; + LOG(INFO) << "[EP] RoCE active QPs/rank override = " + << (active_qps_cap_ > 0 ? std::to_string(active_qps_cap_) + : "auto"); // Get ranks CUDA_CHECK(cudaGetDevice(&device_id)); CUDA_CHECK(cudaDeviceGetAttribute(&clock_rate_khz, cudaDevAttrClockRate, device_id)); - // P2P transport — owns GDR buffer allocation and IPC handle exchange. + // P2P transport owns GDR buffer allocation. Peer mappings remain disabled + // when the EP caller selects RDMA-only operation. if (engine) { p2p_transport_ = engine->getOrCreateP2pTransport(num_ranks); } else { @@ -92,7 +125,7 @@ MooncakeEpBuffer::MooncakeEpBuffer(int rank, int num_ranks, if (rdma_transport_) { if (!initRdmaTransport(rdma_transport_, gdr_buffer, num_ep_buffer_bytes, num_ranks, USE_QP_COUNT, - comm_stream.stream())) { + comm_stream)) { rdma_transport_ = nullptr; ibgda_disabled_ = true; LOG(INFO) << "[EP] IBGDA unavailable, using P2P-only path"; @@ -114,7 +147,7 @@ MooncakeEpBuffer::MooncakeEpBuffer(int rank, int num_ranks, } auto t = device::createIbgdaDeviceTransport(device_filter); if (initRdmaTransport(t.get(), gdr_buffer, num_ep_buffer_bytes, - num_ranks, USE_QP_COUNT, comm_stream.stream())) { + num_ranks, USE_QP_COUNT, comm_stream)) { owned_rdma_transport_ = std::move(t); rdma_transport_ = owned_rdma_transport_.get(); } else { @@ -151,33 +184,28 @@ MooncakeEpBuffer::~MooncakeEpBuffer() noexcept(false) { p2p_transport_ = nullptr; if (workspace) cudaFree(workspace); + if (comm_stream) { + cudaStreamDestroy(comm_stream); + comm_stream = nullptr; + } } -std::tuple, torch::Tensor, - torch::Tensor, torch::Tensor, std::optional, - std::optional>> -MooncakeEpBuffer::dispatch(const torch::Tensor& x, - const torch::Tensor& topk_idx, - torch::Tensor& active_ranks, - int num_max_dispatch_tokens_per_rank, - int num_experts, int timeout_us, bool use_fp8, - bool async, bool return_recv_hook) { - // Tensor checks - // By default using `ptp128c` FP8 cast - EP_HOST_ASSERT(x.dim() == 2 and x.is_contiguous() and - x.scalar_type() == torch::kBFloat16); - EP_HOST_ASSERT(x.size(1) % sizeof(int4) == 0 and x.size(1) % 128 == 0); - EP_HOST_ASSERT(topk_idx.dim() == 2 and topk_idx.is_contiguous()); - EP_HOST_ASSERT(x.size(0) == topk_idx.size(0) and - x.size(0) <= num_max_dispatch_tokens_per_rank); - EP_HOST_ASSERT(topk_idx.scalar_type() == torch::kInt64); +std::tuple, std::optional>> +MooncakeEpBuffer::dispatch( + uint64_t x_ptr, uint64_t topk_idx_ptr, uint64_t active_ranks_ptr, + int num_tokens, int hidden, int num_topk, + int num_max_dispatch_tokens_per_rank, int num_experts, int timeout_us, + bool use_fp8, uint64_t packed_recv_x_ptr, uint64_t packed_recv_x_scales_ptr, + uint64_t packed_recv_count_ptr, uint64_t packed_recv_src_info_ptr, + uint64_t packed_recv_layout_range_ptr, bool async, bool return_recv_hook, + uint64_t compute_stream_ptr) { EP_HOST_ASSERT(num_experts % num_ranks == 0); EP_HOST_ASSERT(USE_QP_COUNT % num_ranks == 0); + EP_HOST_ASSERT(hidden % static_cast(sizeof(int4)) == 0 && + hidden % 128 == 0); + EP_HOST_ASSERT(num_tokens <= num_max_dispatch_tokens_per_rank); - auto num_tokens = static_cast(x.size(0)), - hidden = static_cast(x.size(1)); - auto num_scales = hidden / 128, - num_topk = static_cast(topk_idx.size(1)); + auto num_scales = hidden / 128; int num_local_experts = num_experts / num_ranks; // Buffer control @@ -190,40 +218,36 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, int phase_epoch = ++phase_epochs[current_buffer_idx]; // Wait previous tasks to be finished - // NOTES: the hook mode will always use the default stream - auto compute_stream = at::cuda::getCurrentCUDAStream(); - auto launch_stream = return_recv_hook ? compute_stream : comm_stream; + // NOTES: the hook mode will always use the default stream, whose native + // handle is allowed to be nullptr in CUDA/PyTorch. + auto compute_stream_raw = + reinterpret_cast(compute_stream_ptr); + auto launch_stream = return_recv_hook ? compute_stream_raw : comm_stream; EP_HOST_ASSERT(not(async and return_recv_hook)); - if (not return_recv_hook) stream_wait(launch_stream, compute_stream); + if (not return_recv_hook) stream_wait(launch_stream, compute_stream_raw); // Allocate packed tensors - auto packed_recv_x = torch::empty( - {num_local_experts, num_ranks * num_max_dispatch_tokens_per_rank, - hidden}, - x.options().dtype(use_fp8 ? torch::kFloat8_e4m3fn : torch::kBFloat16)); - auto packed_recv_src_info = torch::empty( - {num_local_experts, num_ranks * num_max_dispatch_tokens_per_rank}, - torch::dtype(torch::kInt32).device(x.device())); - auto packed_recv_layout_range = - torch::empty({num_local_experts, num_ranks}, - torch::dtype(torch::kInt64).device(x.device())); - auto packed_recv_count = torch::zeros( - {num_local_experts}, torch::dtype(torch::kInt32).device(x.device())); - - // Allocate column-majored scales - auto packed_recv_x_scales = std::optional(); - float* packed_recv_x_scales_ptr = nullptr; + void* x = reinterpret_cast(x_ptr); + auto* topk_idx = reinterpret_cast(topk_idx_ptr); + auto* active_ranks = reinterpret_cast(active_ranks_ptr); + void* packed_recv_x = reinterpret_cast(packed_recv_x_ptr); + auto* packed_recv_x_scales = + reinterpret_cast(packed_recv_x_scales_ptr); + auto* packed_recv_count = reinterpret_cast(packed_recv_count_ptr); + auto* packed_recv_src_info = + reinterpret_cast(packed_recv_src_info_ptr); + auto* packed_recv_layout_range = + reinterpret_cast(packed_recv_layout_range_ptr); + EP_HOST_ASSERT(active_ranks != nullptr); + EP_HOST_ASSERT(num_tokens == 0 || (x != nullptr && topk_idx != nullptr)); + EP_HOST_ASSERT(packed_recv_x != nullptr && packed_recv_count != nullptr); + EP_HOST_ASSERT(packed_recv_src_info != nullptr && + packed_recv_layout_range != nullptr); if (use_fp8) { EP_HOST_ASSERT((num_ranks * num_max_dispatch_tokens_per_rank) % 4 == 0 and "TMA requires the number of tokens to be multiple of 4"); - packed_recv_x_scales = - torch::empty({num_local_experts, num_scales, - num_ranks * num_max_dispatch_tokens_per_rank}, - torch::dtype(torch::kFloat32).device(x.device())); - packed_recv_x_scales = - torch::transpose(packed_recv_x_scales.value(), 1, 2); - packed_recv_x_scales_ptr = packed_recv_x_scales->data_ptr(); + EP_HOST_ASSERT(packed_recv_x_scales != nullptr); } int64_t timeout_ticks = @@ -238,10 +262,9 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, void** ipc_ptrs = p2p_transport_->peerPtrsTablePtr(); int active_qps_per_rank = active_qps_per_rank_for_ep( USE_QP_COUNT / num_ranks, rdma_transport_ && rdma_transport_->isRoce(), - active_qps_cap_); - + active_qps_cap_, num_experts / num_ranks); auto mark_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::mark_phase_ack(gdr_buffer, nvlink_avail, ipc_ptrs, buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream); @@ -249,7 +272,7 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, }; auto wait_peer_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::wait_phase_ack(buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream, timeout_ticks); @@ -257,7 +280,7 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, }; auto mark_and_wait_peer_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::mark_and_wait_phase_ack( gdr_buffer, nvlink_avail, ipc_ptrs, buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream, timeout_ticks); @@ -266,18 +289,16 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, auto launcher = [=](int phases) { mooncake::dispatch( - packed_recv_x.data_ptr(), packed_recv_x_scales_ptr, - packed_recv_src_info.data_ptr(), - packed_recv_layout_range.data_ptr(), - packed_recv_count.data_ptr(), active_ranks.data_ptr(), + packed_recv_x, packed_recv_x_scales, packed_recv_src_info, + packed_recv_layout_range, packed_recv_count, active_ranks, gdr_buffer, buffer.rdma_send_signal_buffer, buffer.rdma_recv_signal_buffer, buffer.rdma_send_data_buffer, buffer.rdma_recv_data_buffer, nullptr, nullptr, raddrs_ptr, - rkeys_ptr, qp_devctxs_ptr, nvlink_avail, ipc_ptrs, x.data_ptr(), - topk_idx.data_ptr(), next_buffer.rdma_recv_signal_buffer, - num_tokens, hidden, num_max_dispatch_tokens_per_rank, num_topk, - num_experts, rank, num_ranks, use_fp8, workspace, launch_stream, - timeout_ticks, phases, active_qps_per_rank); + rkeys_ptr, qp_devctxs_ptr, nvlink_avail, ipc_ptrs, x, topk_idx, + next_buffer.rdma_recv_signal_buffer, num_tokens, hidden, + num_max_dispatch_tokens_per_rank, num_topk, num_experts, rank, + num_ranks, use_fp8, workspace, launch_stream, timeout_ticks, phases, + active_qps_per_rank); }; if (return_recv_hook) { launcher(LOW_LATENCY_SEND_PHASE); @@ -298,11 +319,11 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, // NOTES: we must ensure the all tensors will not be deallocated // before the stream-wait happens, so in Python API, we must wrap // all tensors into the event handle. - event = EventHandle(launch_stream); + event = EventHandle(reinterpret_cast(launch_stream)); } else if (return_recv_hook && macaHostPhaseFenceCoversPeers()) { - event = EventHandle(launch_stream); + event = EventHandle(reinterpret_cast(launch_stream)); } else if (not return_recv_hook) { - stream_wait(compute_stream, launch_stream); + stream_wait(compute_stream_raw, launch_stream); } // Receiver callback @@ -314,50 +335,36 @@ MooncakeEpBuffer::dispatch(const torch::Tensor& x, }; // Return values - return {packed_recv_x, - packed_recv_x_scales, - packed_recv_count, - packed_recv_src_info, - packed_recv_layout_range, - event, - recv_hook}; + return {event, recv_hook}; } -std::tuple, - std::optional>> -MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, - const torch::Tensor& topk_weights, - const torch::Tensor& src_info, - const torch::Tensor& layout_range, - torch::Tensor& active_ranks, +std::tuple, std::optional>> +MooncakeEpBuffer::combine(uint64_t x_ptr, uint64_t topk_idx_ptr, + uint64_t topk_weights_ptr, uint64_t src_info_ptr, + uint64_t layout_range_ptr, uint64_t active_ranks_ptr, + int num_local_experts, int num_combined_tokens, + int hidden, int num_topk, int num_max_dispatch_tokens_per_rank, int num_experts, - int timeout_us, bool zero_copy, bool async, - bool return_recv_hook, - const std::optional& out) { - // Tensor checks - EP_HOST_ASSERT(x.dim() == 3 and x.is_contiguous() and - x.scalar_type() == torch::kBFloat16); - EP_HOST_ASSERT(x.size(0) == num_experts / num_ranks); - EP_HOST_ASSERT(x.size(1) == num_ranks * num_max_dispatch_tokens_per_rank); - EP_HOST_ASSERT(x.size(2) % sizeof(int4) == 0 and x.size(2) % 128 == 0); - EP_HOST_ASSERT(topk_idx.dim() == 2 and topk_idx.is_contiguous()); - EP_HOST_ASSERT(topk_idx.size(0) == topk_weights.size(0) and - topk_idx.size(1) == topk_weights.size(1)); - EP_HOST_ASSERT(topk_idx.scalar_type() == torch::kInt64); - EP_HOST_ASSERT(topk_weights.dim() == 2 and topk_weights.is_contiguous()); - EP_HOST_ASSERT(topk_weights.size(0) <= num_max_dispatch_tokens_per_rank); - EP_HOST_ASSERT(topk_weights.scalar_type() == torch::kFloat32); - EP_HOST_ASSERT(src_info.dim() == 2 and src_info.is_contiguous()); - EP_HOST_ASSERT(src_info.scalar_type() == torch::kInt32 and - x.size(0) == src_info.size(0)); - EP_HOST_ASSERT(layout_range.dim() == 2 and layout_range.is_contiguous()); - EP_HOST_ASSERT(layout_range.scalar_type() == torch::kInt64); - EP_HOST_ASSERT(layout_range.size(0) == num_experts / num_ranks and - layout_range.size(1) == num_ranks); - auto hidden = static_cast(x.size(2)); - auto num_local_experts = num_experts / num_ranks, - num_topk = static_cast(topk_weights.size(1)); - auto num_combined_tokens = static_cast(topk_weights.size(0)); + int timeout_us, bool zero_copy, + uint64_t combined_x_ptr, bool async, + bool return_recv_hook, uint64_t compute_stream_ptr) { + EP_HOST_ASSERT(num_local_experts == num_experts / num_ranks); + EP_HOST_ASSERT(hidden % static_cast(sizeof(int4)) == 0 && + hidden % 128 == 0); + EP_HOST_ASSERT(num_combined_tokens <= num_max_dispatch_tokens_per_rank); + void* x = reinterpret_cast(x_ptr); + auto* topk_idx = reinterpret_cast(topk_idx_ptr); + auto* topk_weights = reinterpret_cast(topk_weights_ptr); + auto* src_info = reinterpret_cast(src_info_ptr); + auto* layout_range = reinterpret_cast(layout_range_ptr); + auto* active_ranks = reinterpret_cast(active_ranks_ptr); + void* combined_x = reinterpret_cast(combined_x_ptr); + EP_HOST_ASSERT( + num_combined_tokens == 0 || + (x != nullptr && topk_idx != nullptr && topk_weights != nullptr)); + EP_HOST_ASSERT(src_info != nullptr && layout_range != nullptr); + EP_HOST_ASSERT(active_ranks != nullptr); + EP_HOST_ASSERT(num_combined_tokens == 0 || combined_x != nullptr); // Buffer control BufferPair layout(gdr_buffer, num_max_dispatch_tokens_per_rank, hidden, @@ -369,23 +376,13 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, int phase_epoch = ++phase_epochs[current_buffer_idx]; // Wait previous tasks to be finished - // NOTES: the hook mode will always use the default stream - auto compute_stream = at::cuda::getCurrentCUDAStream(); - auto launch_stream = return_recv_hook ? compute_stream : comm_stream; + // NOTES: the hook mode will always use the default stream, whose native + // handle is allowed to be nullptr in CUDA/PyTorch. + auto compute_stream_raw = + reinterpret_cast(compute_stream_ptr); + auto launch_stream = return_recv_hook ? compute_stream_raw : comm_stream; EP_HOST_ASSERT(not(async and return_recv_hook)); - if (not return_recv_hook) stream_wait(launch_stream, compute_stream); - - // Allocate output tensor - torch::Tensor combined_x; - if (out.has_value()) { - EP_HOST_ASSERT(out->dim() == 2 and out->is_contiguous()); - EP_HOST_ASSERT(out->size(0) == num_combined_tokens and - out->size(1) == hidden); - EP_HOST_ASSERT(out->scalar_type() == x.scalar_type()); - combined_x = out.value(); - } else { - combined_x = torch::empty({num_combined_tokens, hidden}, x.options()); - } + if (not return_recv_hook) stream_wait(launch_stream, compute_stream_raw); int64_t timeout_ticks = timeout_us == -1 ? -1 @@ -399,10 +396,10 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, void** ipc_ptrs = p2p_transport_->peerPtrsTablePtr(); int active_qps_per_rank = active_qps_per_rank_for_ep( USE_QP_COUNT / num_ranks, rdma_transport_ && rdma_transport_->isRoce(), - active_qps_cap_); + active_qps_cap_, num_experts / num_ranks); auto mark_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::mark_phase_ack(gdr_buffer, nvlink_avail, ipc_ptrs, buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream); @@ -410,7 +407,7 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, }; auto wait_peer_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::wait_phase_ack(buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream, timeout_ticks); @@ -418,7 +415,7 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, }; auto mark_and_wait_peer_send_done = [=]() { -#ifdef MOONCAKE_EP_SPLIT_SEND_RECV +#ifdef MOONCAKE_EP_PHASE_ACK mooncake::mark_and_wait_phase_ack( gdr_buffer, nvlink_avail, ipc_ptrs, buffer.rdma_send_signal_buffer, rank, num_ranks, phase_epoch, launch_stream, timeout_ticks); @@ -428,13 +425,11 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, // Kernel launch auto launcher = [=](int phases) { mooncake::combine( - combined_x.data_ptr(), active_ranks.data_ptr(), gdr_buffer, + combined_x, active_ranks, gdr_buffer, buffer.rdma_send_signal_buffer, buffer.rdma_recv_signal_buffer, buffer.rdma_send_data_buffer, buffer.rdma_recv_data_buffer, nullptr, nullptr, raddrs_ptr, rkeys_ptr, qp_devctxs_ptr, nvlink_avail, - ipc_ptrs, x.data_ptr(), topk_idx.data_ptr(), - topk_weights.data_ptr(), src_info.data_ptr(), - layout_range.data_ptr(), + ipc_ptrs, x, topk_idx, topk_weights, src_info, layout_range, next_buffer.rdma_recv_signal_buffer, num_combined_tokens, hidden, num_max_dispatch_tokens_per_rank, num_topk, num_experts, rank, num_ranks, workspace, launch_stream, timeout_ticks, phases, @@ -459,11 +454,11 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, // NOTES: we must ensure the all tensors will not be deallocated // before the stream-wait happens, so in Python API, we must wrap // all tensors into the event handle. - event = EventHandle(launch_stream); + event = EventHandle(reinterpret_cast(launch_stream)); } else if (return_recv_hook && macaHostPhaseFenceCoversPeers()) { - event = EventHandle(launch_stream); + event = EventHandle(reinterpret_cast(launch_stream)); } else if (not return_recv_hook) { - stream_wait(compute_stream, launch_stream); + stream_wait(compute_stream_raw, launch_stream); } // Receiver callback @@ -475,33 +470,12 @@ MooncakeEpBuffer::combine(const torch::Tensor& x, const torch::Tensor& topk_idx, }; // Return values - return {combined_x, event, recv_hook}; -} - -torch::Tensor MooncakeEpBuffer::get_next_combine_buffer( - int num_max_dispatch_tokens_per_rank, int hidden, int num_experts) { - BufferPair layout(gdr_buffer, num_max_dispatch_tokens_per_rank, hidden, - num_ranks, num_experts); - - auto buffer = layout.buffers[buffer_idx]; - auto dtype = torch::kBFloat16; - size_t num_bytes_per_combine_msg = hidden * EP_BF16_SIZE; - auto num_msg_elems = - static_cast(num_bytes_per_combine_msg / elementSize(dtype)); - - EP_HOST_ASSERT(num_bytes_per_combine_msg % elementSize(dtype) == 0); - return torch::from_blob( - buffer.rdma_send_data_buffer, - {num_experts / num_ranks, num_ranks * num_max_dispatch_tokens_per_rank, - hidden}, - {num_ranks * num_max_dispatch_tokens_per_rank * num_msg_elems, - num_msg_elems, 1}, - torch::TensorOptions().dtype(dtype).device(torch::kCUDA)); + return {event, recv_hook}; } void MooncakeEpBuffer::update_local_qpns() { if (!rdma_transport_) return; - int ret = rdma_transport_->recreateQueuePairs(comm_stream.stream()); + int ret = rdma_transport_->recreateQueuePairs(comm_stream); if (ret != 0) { ibgda_disabled_ = true; LOG(ERROR) << "[EP] Failed to recreate QPs"; @@ -542,12 +516,14 @@ void MooncakeEpBuffer::sync_ibgda_peers( } std::vector MooncakeEpBuffer::get_ipc_handle() { + if (!p2p_enabled_) return {}; return p2p_transport_->exportIpcHandle(gdr_buffer); } void MooncakeEpBuffer::sync_nvlink_ipc_handles( const std::vector>& remote_handles, const std::vector& active_ranks_mask) { + if (!p2p_enabled_) return; p2p_transport_->importPeerHandles(gdr_buffer, rank, num_ranks, remote_handles, active_ranks_mask); } diff --git a/mooncake-ep/src/mooncake_ep_elastic_buffer.cpp b/mooncake-ep/src/mooncake_ep_elastic_buffer.cpp index 7d3d0052cf..0b7977f659 100644 --- a/mooncake-ep/src/mooncake_ep_elastic_buffer.cpp +++ b/mooncake-ep/src/mooncake_ep_elastic_buffer.cpp @@ -6,7 +6,7 @@ #include #include -#include +#include namespace mooncake { namespace { @@ -197,6 +197,23 @@ std::tuple MooncakeElasticBuffer::get_logical_domain_size() const { return {topology_.num_scaleout_ranks, topology_.num_scaleup_ranks}; } +std::shared_ptr +MooncakeElasticBuffer::ensure_deterministic_rank_count_buffer(int num_sms) { + const int64_t required_bytes = static_cast(sizeof(int)) * num_sms * + topology_.num_scaleup_ranks; + if (deterministic_rank_count_buffer_ != nullptr && + deterministic_rank_count_buffer_bytes_ >= required_bytes) { + return deterministic_rank_count_buffer_; + } + + void* buffer_ptr = nullptr; + CUDA_CHECK(cudaMalloc(&buffer_ptr, required_bytes)); + deterministic_rank_count_buffer_ = + std::shared_ptr(buffer_ptr, [](void* p) { cudaFree(p); }); + deterministic_rank_count_buffer_bytes_ = required_bytes; + return deterministic_rank_count_buffer_; +} + int MooncakeElasticBuffer::get_theoretical_num_sms(int num_experts, int num_topk) const { int device = 0; @@ -210,41 +227,26 @@ int MooncakeElasticBuffer::get_theoretical_num_sms(int num_experts, std::max(1, num_experts * num_topk)})); } -ElasticDispatchOutput MooncakeElasticBuffer::dispatch( - const torch::Tensor& x, const std::optional& sf, - const torch::Tensor& topk_idx, - const std::optional& topk_weights, - torch::Tensor& active_ranks, int num_experts, int num_max_tokens_per_rank, - int expert_alignment, int num_sms, bool do_expand, bool do_cpu_sync, - bool async_with_compute_stream, - const std::optional& cached_handle) { - EP_HOST_ASSERT(x.dim() == 2 && x.is_contiguous()); - const bool use_sf = sf.has_value(); - if (use_sf) { - EP_HOST_ASSERT(x.element_size() == 1); - EP_HOST_ASSERT(sf->dim() == 2 && sf->is_cuda()); - EP_HOST_ASSERT(sf->scalar_type() == torch::kFloat32 || - sf->scalar_type() == torch::kInt32); - EP_HOST_ASSERT(sf->size(0) == x.size(0)); - } else { - EP_HOST_ASSERT(!config_.use_fp8_dispatch); - EP_HOST_ASSERT(x.scalar_type() == torch::kBFloat16); - } - EP_HOST_ASSERT(topk_idx.dim() == 2 && topk_idx.is_contiguous()); - EP_HOST_ASSERT(topk_idx.scalar_type() == torch::kInt64); - EP_HOST_ASSERT(x.size(0) == topk_idx.size(0)); +std::optional MooncakeElasticBuffer::dispatch( + uint64_t x_ptr, int x_element_size, uint64_t sf_ptr, int num_tokens, + int hidden, int num_sf_packs, int sf_token_stride, int sf_hidden_stride, + uint64_t topk_idx_ptr, int num_topk, uint64_t topk_weights_ptr, + uint64_t active_ranks_ptr, int num_experts, int num_max_tokens_per_rank, + int expert_alignment, int num_sms, bool do_expand, + bool async_with_compute_stream, uint64_t compute_stream_ptr, + bool cached_mode, uint64_t psum_num_recv_tokens_per_scaleup_rank_ptr, + uint64_t psum_num_recv_tokens_per_expert_ptr, + uint64_t dst_buffer_slot_idx_ptr, uint64_t token_metadata_at_forward_ptr, + uint64_t channel_linked_list_ptr, uint64_t recv_x_ptr, + uint64_t recv_x_scales_ptr, uint64_t recv_topk_idx_ptr, + uint64_t recv_topk_weights_ptr, uint64_t recv_src_metadata_ptr) { + const bool use_sf = sf_ptr != 0; EP_HOST_ASSERT(num_experts % topology_.num_ranks == 0); - const int num_tokens = static_cast(x.size(0)); - const int hidden = static_cast(x.size(1)); - const int num_topk = static_cast(topk_idx.size(1)); - const int num_sf_packs = use_sf ? static_cast(sf->size(1)) : 0; - const int sf_token_stride = use_sf ? static_cast(sf->stride(0)) : 0; - const int sf_hidden_stride = use_sf ? static_cast(sf->stride(1)) : 0; const int num_local_experts = num_experts / topology_.num_ranks; // The copy epilogue uses `kNumMaxTokensPerRank * kNumRanks` as the // no-CPU-sync sentinel and then reads the real local receive count from the - // GPU prefix-sum tensor. In hybrid mode each scale-up peer may receive + // GPU prefix-sum tensor. In hybrid mode each scale-up peer may receive // tokens forwarded from every scale-out rank, so the conservative output // capacity and sentinel must cover the full logical world, not just the // intra-node scale-up domain. @@ -252,42 +254,47 @@ ElasticDispatchOutput MooncakeElasticBuffer::dispatch( const int num_smem_bytes = device_smem_bytes(); const int num_channels_per_sm = 1; const int num_channels = num_sms * num_channels_per_sm; - const bool cached_mode = cached_handle.has_value(); const bool use_hybrid = topology_.num_scaleout_ranks != 1; const int hybrid_channels = use_hybrid ? hybrid_num_channels(num_sms) : 0; const int hybrid_max_tokens_per_channel = use_hybrid ? hybrid_num_max_tokens_per_channel(num_max_tokens_per_rank, num_sms) : 0; - if (cached_mode) { - const auto& handle = cached_handle.value(); - EP_HOST_ASSERT(!handle.do_expand && !do_expand); - EP_HOST_ASSERT(handle.num_experts == num_experts); - EP_HOST_ASSERT(handle.expert_alignment == expert_alignment); - EP_HOST_ASSERT(handle.num_max_tokens_per_rank == - num_max_tokens_per_rank); - EP_HOST_ASSERT(handle.num_sms == num_sms); - if (use_hybrid) { - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.dim() == 4); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(0) == - hybrid_channels); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(1) == - topology_.num_scaleout_ranks); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(2) == - hybrid_max_tokens_per_channel); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(3) == num_topk); - EP_HOST_ASSERT(handle.token_metadata_at_forward.has_value()); - EP_HOST_ASSERT(handle.channel_linked_list.has_value()); - } else { - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.dim() == 2); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(0) == num_tokens); - EP_HOST_ASSERT(handle.dst_buffer_slot_idx.size(1) == num_topk); - } + + EP_HOST_ASSERT(x_ptr != 0 && topk_idx_ptr != 0 && active_ranks_ptr != 0); + EP_HOST_ASSERT(psum_num_recv_tokens_per_scaleup_rank_ptr != 0); + EP_HOST_ASSERT(psum_num_recv_tokens_per_expert_ptr != 0); + EP_HOST_ASSERT(dst_buffer_slot_idx_ptr != 0); + EP_HOST_ASSERT(recv_x_ptr != 0 && recv_topk_idx_ptr != 0 && + recv_src_metadata_ptr != 0); + if (use_hybrid) { + EP_HOST_ASSERT(token_metadata_at_forward_ptr != 0); + EP_HOST_ASSERT(channel_linked_list_ptr != 0); } - auto compute_stream = at::cuda::getCurrentCUDAStream(); + void* x = reinterpret_cast(x_ptr); + void* sf = reinterpret_cast(sf_ptr); + auto* topk_idx = reinterpret_cast(topk_idx_ptr); + auto* topk_weights = reinterpret_cast(topk_weights_ptr); + auto* active_ranks = reinterpret_cast(active_ranks_ptr); + auto* psum_num_recv_tokens_per_scaleup_rank = + reinterpret_cast(psum_num_recv_tokens_per_scaleup_rank_ptr); + auto* psum_num_recv_tokens_per_expert = + reinterpret_cast(psum_num_recv_tokens_per_expert_ptr); + auto* dst_buffer_slot_idx = reinterpret_cast(dst_buffer_slot_idx_ptr); + auto* token_metadata_at_forward = + reinterpret_cast(token_metadata_at_forward_ptr); + auto* channel_linked_list = reinterpret_cast(channel_linked_list_ptr); + void* recv_x = reinterpret_cast(recv_x_ptr); + void* recv_x_scales = reinterpret_cast(recv_x_scales_ptr); + auto* recv_topk_idx = reinterpret_cast(recv_topk_idx_ptr); + auto* recv_topk_weights = reinterpret_cast(recv_topk_weights_ptr); + auto* recv_src_metadata = reinterpret_cast(recv_src_metadata_ptr); + + auto compute_stream_raw = + reinterpret_cast(compute_stream_ptr); auto launch_stream = native_buffer_->comm_stream; - stream_wait(launch_stream, compute_stream); + stream_wait(launch_stream, compute_stream_raw); const int64_t timeout_cycles = config_.num_gpu_timeout_secs < 0 @@ -297,55 +304,7 @@ ElasticDispatchOutput MooncakeElasticBuffer::dispatch( auto launch_ctx = make_launch_context( *native_buffer_, topology_, mapped_host_workspace_, timeout_cycles); - auto psum_num_recv_tokens_per_scaleup_rank = - cached_mode ? cached_handle->psum_num_recv_tokens_per_scaleup_rank - : torch::empty({topology_.num_scaleup_ranks}, - torch::TensorOptions() - .dtype(torch::kInt32) - .device(x.device())); - auto psum_num_recv_tokens_per_expert = - cached_mode - ? cached_handle->psum_num_recv_tokens_per_expert - : torch::empty({num_local_experts + 1}, torch::TensorOptions() - .dtype(torch::kInt32) - .device(x.device())); - auto dst_buffer_slot_idx = - cached_mode - ? cached_handle->dst_buffer_slot_idx - : (use_hybrid ? torch::empty( - {hybrid_channels, topology_.num_scaleout_ranks, - hybrid_max_tokens_per_channel, num_topk}, - torch::TensorOptions() - .dtype(torch::kInt32) - .device(x.device())) - : torch::empty({num_tokens, num_topk}, - torch::TensorOptions() - .dtype(torch::kInt32) - .device(x.device()))); - std::optional token_metadata_at_forward = std::nullopt; - std::optional channel_linked_list = std::nullopt; - if (use_hybrid) { - if (cached_mode) { - token_metadata_at_forward = - cached_handle->token_metadata_at_forward; - channel_linked_list = cached_handle->channel_linked_list; - } else { - const int forward_metadata_dims = 2 + num_topk * 2; - token_metadata_at_forward = torch::empty( - {hybrid_channels, - topology_.num_scaleout_ranks * hybrid_max_tokens_per_channel + - 1, - forward_metadata_dims}, - torch::TensorOptions().dtype(torch::kInt32).device(x.device())); - channel_linked_list = torch::empty( - {hybrid_channels, - topology_.num_scaleout_ranks * hybrid_max_tokens_per_channel + - 1, - topology_.num_scaleup_ranks}, - torch::TensorOptions().dtype(torch::kInt32).device(x.device())); - } - } - std::optional deterministic_rank_count_buffer = std::nullopt; + std::shared_ptr deterministic_rank_count_buffer; #ifdef MOONCAKE_EP_USE_MUSA // MUSA non-hybrid dispatch always runs // launch_musa_elastic_prepare_dispatch(), which assigns slots and publishes @@ -356,200 +315,87 @@ ElasticDispatchOutput MooncakeElasticBuffer::dispatch( config_.deterministic && !cached_mode && !use_hybrid; #endif if (run_deterministic_prologue) { - deterministic_rank_count_buffer = torch::empty( - {num_sms, topology_.num_scaleup_ranks}, - torch::TensorOptions().dtype(torch::kInt32).device(x.device())); + deterministic_rank_count_buffer = + ensure_deterministic_rank_count_buffer(num_sms); launch_elastic_dispatch_deterministic_prologue( - topk_idx.data_ptr(), - deterministic_rank_count_buffer.value().data_ptr(), - dst_buffer_slot_idx.data_ptr(), num_tokens, - num_max_tokens_per_rank, num_experts, num_topk, - topology_.scaleup_rank_idx, topology_.num_scaleup_ranks, num_sms, - num_smem_bytes, launch_stream.stream()); + topk_idx, static_cast(deterministic_rank_count_buffer.get()), + dst_buffer_slot_idx, num_tokens, num_max_tokens_per_rank, + num_experts, num_topk, topology_.scaleup_rank_idx, + topology_.num_scaleup_ranks, num_sms, num_smem_bytes, + launch_stream); } launch_mooncake_elastic_dispatch( - x.data_ptr(), use_sf ? const_cast(sf->data_ptr()) : nullptr, - const_cast(topk_idx.data_ptr()), - topk_weights.has_value() - ? const_cast(topk_weights->data_ptr()) - : nullptr, - nullptr, nullptr, psum_num_recv_tokens_per_scaleup_rank.data_ptr(), - psum_num_recv_tokens_per_expert.data_ptr(), - dst_buffer_slot_idx.data_ptr(), - token_metadata_at_forward.has_value() - ? token_metadata_at_forward->data_ptr() - : nullptr, - num_tokens, num_max_tokens_per_rank, hidden, - static_cast(x.element_size()), num_sf_packs, sf_token_stride, - sf_hidden_stride, num_experts, num_topk, expert_alignment, num_sms, + x, sf, topk_idx, topk_weights, nullptr, nullptr, + psum_num_recv_tokens_per_scaleup_rank, psum_num_recv_tokens_per_expert, + dst_buffer_slot_idx, token_metadata_at_forward, num_tokens, + num_max_tokens_per_rank, hidden, x_element_size, num_sf_packs, + sf_token_stride, sf_hidden_stride, num_experts, num_topk, + expert_alignment, num_sms, use_hybrid ? kElasticHybridChannelsPerSm : num_channels_per_sm, num_smem_bytes, cached_mode, config_.deterministic, false, launch_ctx, - launch_stream.stream()); - - const int num_recv_output_capacity = - do_expand ? num_recv_tokens * num_topk : num_recv_tokens; - auto recv_x = torch::empty({num_recv_output_capacity, hidden}, x.options()); - auto recv_x_scales = std::optional(); - void* recv_x_scales_ptr = nullptr; - int recv_sf_token_stride = 0; - int recv_sf_hidden_stride = 0; - if (use_sf) { - recv_x_scales = torch::empty({num_recv_output_capacity, num_sf_packs}, - sf->options()); - recv_x_scales_ptr = recv_x_scales->data_ptr(); - recv_sf_token_stride = static_cast(recv_x_scales->stride(0)); - recv_sf_hidden_stride = static_cast(recv_x_scales->stride(1)); - } - auto recv_topk_idx = - torch::empty({num_recv_tokens, num_topk}, topk_idx.options()); - auto recv_topk_weights = std::optional(); - float* recv_topk_weights_ptr = nullptr; - if (topk_weights.has_value()) { - recv_topk_weights = do_expand - ? torch::empty({num_recv_output_capacity}, - topk_weights->options()) - : torch::empty({num_recv_tokens, num_topk}, - topk_weights->options()); - recv_topk_weights_ptr = recv_topk_weights->data_ptr(); - } - auto recv_src_metadata = torch::empty( - {num_recv_tokens, num_topk + 2}, - torch::TensorOptions().dtype(torch::kInt32).device(x.device())); - auto handle_psum_num_recv_tokens_per_expert = - do_expand - ? psum_num_recv_tokens_per_expert.slice(0, 0, num_local_experts) - : psum_num_recv_tokens_per_expert.slice(0, 1, - num_local_experts + 1); - auto epilogue_psum_num_recv_tokens_per_expert = + launch_stream); + + const int recv_sf_token_stride = num_sf_packs; + const int recv_sf_hidden_stride = 1; + auto* epilogue_psum_num_recv_tokens_per_expert = do_expand ? psum_num_recv_tokens_per_expert - : handle_psum_num_recv_tokens_per_expert; + : psum_num_recv_tokens_per_expert + 1; launch_mooncake_elastic_dispatch_copy_epilogue( - recv_x.data_ptr(), recv_x_scales_ptr, recv_topk_idx.data_ptr(), - recv_topk_weights_ptr, recv_src_metadata.data_ptr(), - channel_linked_list.has_value() ? channel_linked_list->data_ptr() - : nullptr, - num_recv_tokens, num_max_tokens_per_rank, hidden, - static_cast(x.element_size()), num_sf_packs, recv_sf_token_stride, - recv_sf_hidden_stride, num_experts, num_topk, num_sms, num_smem_bytes, - use_hybrid ? hybrid_channels : num_channels, do_expand, cached_mode, - launch_ctx, psum_num_recv_tokens_per_scaleup_rank.data_ptr(), - epilogue_psum_num_recv_tokens_per_expert.data_ptr(), - launch_stream.stream()); - - if (do_cpu_sync || !async_with_compute_stream) { - stream_wait(compute_stream, launch_stream); - } - std::optional event = std::nullopt; - if (async_with_compute_stream) { - event = EventHandle(launch_stream); - } + recv_x, recv_x_scales, recv_topk_idx, recv_topk_weights, + recv_src_metadata, channel_linked_list, num_recv_tokens, + num_max_tokens_per_rank, hidden, x_element_size, num_sf_packs, + recv_sf_token_stride, recv_sf_hidden_stride, num_experts, num_topk, + num_sms, num_smem_bytes, use_hybrid ? hybrid_channels : num_channels, + do_expand, cached_mode, launch_ctx, + psum_num_recv_tokens_per_scaleup_rank, + epilogue_psum_num_recv_tokens_per_expert, launch_stream); - std::vector num_recv_tokens_per_expert_list; - int actual_num_recv_tokens = num_recv_tokens; - int actual_num_output_tokens = num_recv_tokens; - if (do_cpu_sync) { - auto scaleup_psum_cpu = psum_num_recv_tokens_per_scaleup_rank.cpu(); - auto expert_psum_cpu = psum_num_recv_tokens_per_expert.cpu(); - const auto* scaleup_psum = scaleup_psum_cpu.data_ptr(); - const auto* expert_psum = expert_psum_cpu.data_ptr(); - actual_num_recv_tokens = scaleup_psum[topology_.num_scaleup_ranks - 1]; - EP_HOST_ASSERT(actual_num_recv_tokens >= 0 && - actual_num_recv_tokens <= num_recv_tokens); - actual_num_output_tokens = actual_num_recv_tokens; - - num_recv_tokens_per_expert_list.reserve(num_local_experts); - const auto align_count = [expert_alignment](int value) { - return ((value + expert_alignment - 1) / expert_alignment) * - expert_alignment; - }; - if (do_expand) { - int previous_psum = 0; - for (int i = 0; i < num_local_experts; ++i) { - const int count = expert_psum[i] - align_count(previous_psum); - EP_HOST_ASSERT(count >= 0); - num_recv_tokens_per_expert_list.push_back(count); - previous_psum = expert_psum[i]; - } - actual_num_output_tokens = - num_local_experts == 0 ? 0 : expert_psum[num_local_experts - 1]; - } else { - for (int i = 0; i < num_local_experts; ++i) { - const int count = expert_psum[i + 1] - expert_psum[i]; - EP_HOST_ASSERT(count >= 0); - num_recv_tokens_per_expert_list.push_back(count); - } - } - EP_HOST_ASSERT(actual_num_output_tokens >= 0 && - actual_num_output_tokens <= recv_x.size(0)); - - recv_x = recv_x.slice(0, 0, actual_num_output_tokens); - if (recv_x_scales.has_value()) { - recv_x_scales = - recv_x_scales->slice(0, 0, actual_num_output_tokens); - } - recv_topk_idx = recv_topk_idx.slice(0, 0, actual_num_recv_tokens); - if (recv_topk_weights.has_value()) { - recv_topk_weights = - recv_topk_weights->slice(0, 0, actual_num_output_tokens); - } - recv_src_metadata = - recv_src_metadata.slice(0, 0, actual_num_recv_tokens); + (void)active_ranks; + (void)num_local_experts; + (void)hybrid_max_tokens_per_channel; + if (!async_with_compute_stream) { + stream_wait(compute_stream_raw, launch_stream); + return std::nullopt; } - - ElasticNativeHandle handle; - handle.do_expand = do_expand; - handle.num_experts = num_experts; - handle.expert_alignment = expert_alignment; - handle.num_max_tokens_per_rank = num_max_tokens_per_rank; - handle.num_sms = num_sms; - handle.topk_idx = cached_mode ? cached_handle->topk_idx : topk_idx.clone(); - handle.psum_num_recv_tokens_per_expert = - handle_psum_num_recv_tokens_per_expert; - handle.psum_num_recv_tokens_per_scaleup_rank = - psum_num_recv_tokens_per_scaleup_rank; - handle.recv_src_metadata = recv_src_metadata; - handle.recv_layout_range = torch::empty( - {0}, torch::TensorOptions().dtype(torch::kInt64).device(x.device())); - handle.dst_buffer_slot_idx = dst_buffer_slot_idx; - handle.token_metadata_at_forward = token_metadata_at_forward; - handle.channel_linked_list = channel_linked_list; - handle.num_recv_tokens_per_expert_list = num_recv_tokens_per_expert_list; - - ElasticDispatchOutput output; - output.recv_x = recv_x; - output.recv_x_scales = recv_x_scales; - output.recv_topk_idx = recv_topk_idx; - output.recv_topk_weights = recv_topk_weights; - output.handle = handle; - output.event = event; - return output; + return EventHandle(reinterpret_cast(launch_stream), + deterministic_rank_count_buffer); } -ElasticCombineOutput MooncakeElasticBuffer::combine( - const torch::Tensor& x, const ElasticNativeHandle& handle, - const std::optional& topk_weights, - torch::Tensor& active_ranks, int num_sms, bool async_with_compute_stream, - const std::optional& out) { - EP_HOST_ASSERT(x.dim() == 2 && x.is_contiguous()); - EP_HOST_ASSERT(x.scalar_type() == torch::kBFloat16); - torch::Tensor weights = topk_weights.value_or(torch::Tensor()); - if (!weights.defined()) { - weights = torch::ones( - handle.topk_idx.sizes(), - torch::TensorOptions().dtype(torch::kFloat32).device(x.device())); - } - const int hidden = static_cast(x.size(1)); - const int num_topk = static_cast(handle.topk_idx.size(1)); - const int num_combined_tokens = static_cast(handle.topk_idx.size(0)); +std::optional MooncakeElasticBuffer::combine( + uint64_t x_ptr, int num_input_tokens, int hidden, uint64_t topk_idx_ptr, + int num_combined_tokens, int num_topk, uint64_t topk_weights_ptr, + uint64_t psum_num_recv_tokens_per_scaleup_rank_ptr, + uint64_t recv_src_metadata_ptr, uint64_t token_metadata_at_forward_ptr, + uint64_t channel_linked_list_ptr, uint64_t active_ranks_ptr, + int num_experts, int num_max_tokens_per_rank, bool do_expand, int num_sms, + bool async_with_compute_stream, uint64_t compute_stream_ptr, + uint64_t combined_x_ptr) { + EP_HOST_ASSERT(x_ptr != 0 && topk_idx_ptr != 0 && topk_weights_ptr != 0); + EP_HOST_ASSERT(psum_num_recv_tokens_per_scaleup_rank_ptr != 0); + EP_HOST_ASSERT(recv_src_metadata_ptr != 0 && active_ranks_ptr != 0); + EP_HOST_ASSERT(combined_x_ptr != 0); + void* x = reinterpret_cast(x_ptr); + auto* topk_idx = reinterpret_cast(topk_idx_ptr); + auto* topk_weights = reinterpret_cast(topk_weights_ptr); + auto* psum_num_recv_tokens_per_scaleup_rank = + reinterpret_cast(psum_num_recv_tokens_per_scaleup_rank_ptr); + auto* recv_src_metadata = reinterpret_cast(recv_src_metadata_ptr); + auto* token_metadata_at_forward = + reinterpret_cast(token_metadata_at_forward_ptr); + auto* channel_linked_list = reinterpret_cast(channel_linked_list_ptr); + auto* active_ranks = reinterpret_cast(active_ranks_ptr); + void* combined_x = reinterpret_cast(combined_x_ptr); + const int num_smem_bytes = device_smem_bytes(); const int num_channels = std::max(1, num_sms); const bool use_hybrid = topology_.num_scaleout_ranks != 1; const int hybrid_channels = use_hybrid ? hybrid_num_channels(num_sms) : 0; - auto compute_stream = at::cuda::getCurrentCUDAStream(); + auto compute_stream_raw = + reinterpret_cast(compute_stream_ptr); auto launch_stream = native_buffer_->comm_stream; - stream_wait(launch_stream, compute_stream); + stream_wait(launch_stream, compute_stream_raw); const int64_t timeout_cycles = config_.num_gpu_timeout_secs < 0 ? -1 @@ -557,47 +403,26 @@ ElasticCombineOutput MooncakeElasticBuffer::combine( static_cast(config_.num_gpu_timeout_secs) * 1000; auto launch_ctx = make_launch_context( *native_buffer_, topology_, mapped_host_workspace_, timeout_cycles); - auto psum_num_recv_tokens_per_scaleup_rank = - handle.psum_num_recv_tokens_per_scaleup_rank; void* reduce_buffer = launch_mooncake_elastic_combine( - x.data_ptr(), weights.data_ptr(), - const_cast(handle.recv_src_metadata.data_ptr()), - psum_num_recv_tokens_per_scaleup_rank.data_ptr(), - handle.token_metadata_at_forward.has_value() - ? handle.token_metadata_at_forward->data_ptr() - : nullptr, - handle.channel_linked_list.has_value() - ? handle.channel_linked_list->data_ptr() - : nullptr, - static_cast(x.size(0)), handle.num_max_tokens_per_rank, hidden, - handle.num_experts, num_topk, num_sms, num_smem_bytes, - use_hybrid ? hybrid_channels : num_channels, handle.do_expand, - config_.allow_multiple_reduction, launch_ctx, launch_stream.stream()); - - torch::Tensor combined_x = - out.has_value() - ? out.value() - : torch::empty({num_combined_tokens, hidden}, x.options()); + x, topk_weights, recv_src_metadata, + psum_num_recv_tokens_per_scaleup_rank, token_metadata_at_forward, + channel_linked_list, num_input_tokens, num_max_tokens_per_rank, hidden, + num_experts, num_topk, num_sms, num_smem_bytes, + use_hybrid ? hybrid_channels : num_channels, do_expand, + config_.allow_multiple_reduction, launch_ctx, launch_stream); + launch_mooncake_elastic_combine_reduce_epilogue( - combined_x.data_ptr(), weights.data_ptr(), - const_cast(handle.topk_idx.data_ptr()), - num_combined_tokens, handle.num_max_tokens_per_rank, hidden, - handle.num_experts, num_topk, reduce_buffer, nullptr, nullptr, num_sms, - num_smem_bytes, handle.do_expand, config_.allow_multiple_reduction, - launch_ctx, launch_stream.stream()); + combined_x, topk_weights, topk_idx, num_combined_tokens, + num_max_tokens_per_rank, hidden, num_experts, num_topk, reduce_buffer, + nullptr, nullptr, num_sms, num_smem_bytes, do_expand, + config_.allow_multiple_reduction, launch_ctx, launch_stream); + (void)active_ranks; if (!async_with_compute_stream) { - stream_wait(compute_stream, launch_stream); + stream_wait(compute_stream_raw, launch_stream); + return std::nullopt; } - std::optional event = std::nullopt; - if (async_with_compute_stream) event = EventHandle(launch_stream); - (void)active_ranks; - - ElasticCombineOutput output; - output.combined_x = combined_x; - output.combined_topk_weights = std::nullopt; - output.event = event; - return output; + return EventHandle(reinterpret_cast(launch_stream)); } ElasticTopology MooncakeElasticBuffer::discover_topology( diff --git a/mooncake-ep/src/mooncake_ep_kernel.cu b/mooncake-ep/src/mooncake_ep_kernel.cu index bc96441181..47a6af2225 100644 --- a/mooncake-ep/src/mooncake_ep_kernel.cu +++ b/mooncake-ep/src/mooncake_ep_kernel.cu @@ -162,7 +162,7 @@ dispatch(void* packed_recv_x, float* packed_recv_x_scales, const auto warp_group_id = warp_id / kNumWarpsPerGroup; const auto sub_warp_id = warp_id % kNumWarpsPerGroup; const auto responsible_expert_idx = sm_id * kNumWarpGroups + warp_group_id; -#ifdef MOONCAKE_EP_USE_MACA +#if defined(MOONCAKE_EP_USE_MUSA) || defined(MOONCAKE_EP_USE_MACA) // C500 reports 64-thread hardware warps. Do not split the last hardware // warp by assigning only the final 32-thread pseudo-warp to count work. // Reserve one full warp group from the data path, but write counts from a @@ -211,8 +211,8 @@ dispatch(void* packed_recv_x, float* packed_recv_x_scales, // There are 2 kinds of execution lanes in this part: // 1. Data lanes for FP8 cast and sending top-k tokens. // 2. Count lanes for reading `topk_idx` and per-expert token counts. - // MACA reserves a full warp group for the count path; CUDA keeps the - // original final 32-thread warp behavior. + // Non-CUDA backends reserve a full warp group for the count path. This + // keeps the final group out of the data path when MUSA uses five groups. if (is_data_warp) { constexpr int kNumElemsPerRead = sizeof(int4) / EP_BF16_SIZE; EP_DEVICE_ASSERT(kHidden % kNumElemsPerRead == 0); @@ -310,7 +310,8 @@ dispatch(void* packed_recv_x, float* packed_recv_x_scales, // Participate in __syncthreads() barriers from data warps. // Each token iteration in the send loop above calls // __syncthreads() once; the count path must match. - for (int token_idx = sm_id; token_idx < num_tokens; token_idx += num_sms) { + for (int token_idx = sm_id; token_idx < num_tokens; + token_idx += num_sms) { __syncthreads(); } #endif @@ -481,20 +482,26 @@ void dispatch(void* packed_recv_x, float* packed_recv_x_scales, int* next_clean_buffer, int num_tokens, int hidden, int num_max_dispatch_tokens_per_rank, int num_topk, int num_experts, int rank, int num_ranks, bool use_fp8, - void* workspace, cudaStream_t stream, int64_t timeout_ticks, - int phases, int active_qps_per_rank) { + void* workspace, cudaStream_t stream, + int64_t timeout_ticks, int phases, int active_qps_per_rank) { constexpr int kNumMaxTopK = 11; constexpr int kNumWarpsPerGroup = 4; + int num_warp_groups = 8; #ifdef MOONCAKE_EP_USE_MUSA - // MT S5000 benefits from slightly more CTAs while keeping enough warps for top-k<=11. - constexpr int kNumWarpGroups = 5; -#else - constexpr int kNumWarpGroups = 8; + cudaDeviceProp device_prop{}; + int device = 0; + CUDA_CHECK(cudaGetDevice(&device)); + CUDA_CHECK(cudaGetDeviceProperties(&device_prop, device)); + num_warp_groups = cell_div(num_experts, device_prop.multiProcessorCount); + // MUSA keeps four 32-thread pseudo-warps per group. The range is also + // constrained by the count group and the maximum supported CTA shape. + num_warp_groups = max(3, min(8, num_warp_groups)); #endif - EP_STATIC_ASSERT(kNumMaxTopK + 1 <= kNumWarpGroups * kNumWarpsPerGroup, "Too many top-k selections"); + EP_HOST_ASSERT(kNumMaxTopK + 1 <= num_warp_groups * kNumWarpsPerGroup && + "Too many top-k selections"); - const auto num_warps = kNumWarpGroups * kNumWarpsPerGroup; - const auto num_sms = cell_div(num_experts, kNumWarpGroups); + const auto num_warps = num_warp_groups * kNumWarpsPerGroup; + const auto num_sms = max(2, cell_div(num_experts, num_warp_groups)); EP_HOST_ASSERT(num_topk <= kNumMaxTopK); // Workspace checks @@ -502,7 +509,8 @@ void dispatch(void* packed_recv_x, float* packed_recv_x_scales, auto atomic_finish_counter_per_expert = atomic_counter_per_expert + num_experts; EP_HOST_ASSERT(num_experts * sizeof(int) * 2 <= NUM_WORKSPACE_BYTES); -#define DISPATCH_LAUNCH_CASE(hidden) { \ +#define DISPATCH_LAUNCH_GROUP(hidden, groups) case groups: { \ +constexpr int kNumWarpGroups = groups; \ auto dispatch_func = use_fp8 ? dispatch