Skip to content

Commit 09c1330

Browse files
committed
Tighten checksum option flow
1 parent 23d3292 commit 09c1330

14 files changed

Lines changed: 146 additions & 35 deletions

kv_cache_manager/client/pybind/py_client_binding.cc

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -45,12 +45,14 @@ PYBIND11_MODULE(kvcm_py_client, module) {
4545
.value("ER_TRANSFERCLIENT_INIT_ERROR", kvcm::ClientErrorCode::ER_TRANSFERCLIENT_INIT_ERROR)
4646
.value("ER_MANAGERCLIENT_INIT_ERROR", kvcm::ClientErrorCode::ER_MANAGERCLIENT_INIT_ERROR)
4747
.value("ER_CLIENT_NOT_EXISTS", kvcm::ClientErrorCode::ER_CLIENT_NOT_EXISTS)
48+
.value("ER_INIT_CHECK_BUFFER_ERROR", kvcm::ClientErrorCode::ER_INIT_CHECK_BUFFER_ERROR)
4849
.value("ER_SERVICE_NO_STATUS", kvcm::ClientErrorCode::ER_SERVICE_NO_STATUS)
4950
.value("ER_SERVICE_INTERNAL_ERROR", kvcm::ClientErrorCode::ER_SERVICE_INTERNAL_ERROR)
5051
.value("ER_SERVICE_UNSUPPORTED", kvcm::ClientErrorCode::ER_SERVICE_UNSUPPORTED)
5152
.value("ER_SERVICE_INVALID_ARGUMENT", kvcm::ClientErrorCode::ER_SERVICE_INVALID_ARGUMENT)
5253
.value("ER_SERVICE_DUPLICATE_ENTITY", kvcm::ClientErrorCode::ER_SERVICE_DUPLICATE_ENTITY)
5354
.value("ER_SERVICE_INSTANCE_NOT_EXIST", kvcm::ClientErrorCode::ER_SERVICE_INSTANCE_NOT_EXIST)
55+
.value("ER_SERVICE_NOT_LEADER", kvcm::ClientErrorCode::ER_SERVICE_NOT_LEADER)
5456
.value("ER_SDK_TIMEOUT", kvcm::ClientErrorCode::ER_SDK_TIMEOUT)
5557
.value("ER_GETSDK_ERROR", kvcm::ClientErrorCode::ER_GETSDK_ERROR)
5658
.value("ER_CREATESDK_ERROR", kvcm::ClientErrorCode::ER_CREATESDK_ERROR)
@@ -69,6 +71,8 @@ PYBIND11_MODULE(kvcm_py_client, module) {
6971
.value("ER_CUDA_STREAM_SYNCHRONIZE_ERROR", kvcm::ClientErrorCode::ER_CUDA_STREAM_SYNCHRONIZE_ERROR)
7072
.value("ER_CUDA_STREAM_DESTROY_ERROR", kvcm::ClientErrorCode::ER_CUDA_STREAM_DESTROY_ERROR)
7173
.value("ER_CUDA_HOST_REGISTER_ERROR", kvcm::ClientErrorCode::ER_CUDA_HOST_REGISTER_ERROR)
74+
.value("ER_CHECKSUM_MISMATCH", kvcm::ClientErrorCode::ER_CHECKSUM_MISMATCH)
75+
.value("ER_INLINE_HEADER_INVALID", kvcm::ClientErrorCode::ER_INLINE_HEADER_INVALID)
7276
.finalize();
7377

7478
py::native_enum<kvcm::MemoryType>(module, "MemoryType", "enum.Enum")

kv_cache_manager/client/src/internal/stub/grpc_stub.cc

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -409,12 +409,12 @@ ClientErrorCode GrpcStub::FinishWriteCache(const std::string &trace_id,
409409
const std::string write_session_id,
410410
const BlockMask &success_block,
411411
const Locations &locations,
412-
const std::vector<int64_t> &checksums) {
412+
const FinishWriteOptions &options) {
413413
auto stub = GET_AND_CHECK_STUB();
414414
proto::meta::FinishWriteCacheRequest request;
415415
SetCommonInfo(request, trace_id, instance_id);
416416
request.set_write_session_id(write_session_id);
417-
for (auto checksum : checksums) {
417+
for (auto checksum : options.checksums) {
418418
request.add_checksums(checksum);
419419
}
420420
ProtoConvert::BlockMaskToProto(success_block, request.mutable_success_blocks());

kv_cache_manager/client/src/internal/stub/grpc_stub.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,7 @@ class GrpcStub : public Stub {
6363
const std::string write_session_id,
6464
const BlockMask &success_block,
6565
const Locations &locations,
66-
const std::vector<int64_t> &checksums = {}) override;
66+
const FinishWriteOptions &options = FinishWriteOptions{}) override;
6767

6868
ClientErrorCode RemoveCache(const std::string &trace_id,
6969
const std::string &instance_id,

kv_cache_manager/client/src/internal/stub/stub.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -67,7 +67,7 @@ class Stub {
6767
const std::string write_session_id,
6868
const BlockMask &success_block,
6969
const Locations &locations,
70-
const std::vector<int64_t> &checksums = {}) = 0;
70+
const FinishWriteOptions &options = FinishWriteOptions{}) = 0;
7171

7272
virtual ClientErrorCode RemoveCache(const std::string &trace_id,
7373
const std::string &instance_id,

kv_cache_manager/client/src/internal/stub/test/grpc_stub_test.cc

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -494,7 +494,12 @@ TEST_F(GrpcStubTest, TestFinishWriteCacheWithChecksums) {
494494
const std::vector<int64_t> checksums = {0x11, 0x22, 0x33, 0x44};
495495
ASSERT_EQ(ER_OK,
496496
stub_->FinishWriteCache(
497-
"trace3", "instance1", write_session_id, success_block, unrelated_locations, checksums));
497+
"trace3",
498+
"instance1",
499+
write_session_id,
500+
success_block,
501+
unrelated_locations,
502+
FinishWriteOptions::WithChecksums(checksums)));
498503
}
499504
{
500505
auto [success, result] = stub_->GetCacheLocation("trace4",
@@ -525,7 +530,12 @@ TEST_F(GrpcStubTest, TestFinishWriteCacheRejectsChecksumSizeMismatch) {
525530
{
526531
BlockMask success_block = static_cast<size_t>(4);
527532
ASSERT_EQ(ER_SERVICE_INVALID_ARGUMENT,
528-
stub_->FinishWriteCache("trace3", "instance1", write_session_id, success_block, {}, {0x11}));
533+
stub_->FinishWriteCache("trace3",
534+
"instance1",
535+
write_session_id,
536+
success_block,
537+
{},
538+
FinishWriteOptions::WithChecksums(std::vector<int64_t>{0x11})));
529539
}
530540
}
531541

kv_cache_manager/client/src/meta_client_impl.cc

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -182,8 +182,7 @@ ClientErrorCode MetaClientImpl::FinishWrite(const std::string &trace_id,
182182
DebugStringUtil::ToString(locations).c_str(),
183183
options.checksums.size());
184184
const std::string &instance_id = CHECK_INSTANCE_STUB();
185-
return stub_->FinishWriteCache(
186-
trace_id, instance_id, write_session_id, success_block, locations, options.checksums);
185+
return stub_->FinishWriteCache(trace_id, instance_id, write_session_id, success_block, locations, options);
187186
}
188187

189188
ClientErrorCode MetaClientImpl::RemoveCache(const std::string &trace_id,

kv_cache_manager/client/src/transfer_client_impl.cc

Lines changed: 19 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -292,14 +292,27 @@ ClientErrorCode TransferClientImpl::LoadKvCaches(const UriStrVec &uri_str_vec,
292292
#if defined(USING_CUDA) || defined(USING_MUSA)
293293
const auto &trace_info = options.trace_info;
294294
if (is_check_buffer_) {
295-
bool need_print = (trace_info == nullptr) ? true : trace_info->need_print;
296-
std::vector<int64_t> block_checksums;
295+
const bool need_print = (trace_info == nullptr) ? true : trace_info->need_print;
297296
if (need_print) {
298-
auto handle = sdk_buffer_check_pool_->GetCell();
299-
block_checksums = SdkBufferCheckUtil::GetBlocksHash(
300-
block_buffers, handle->d_iovs, handle->d_crcs, handle->h_iovs, max_check_iov_num_, handle->gpu_stream);
297+
auto invalid_it = std::find_if(block_buffers.begin(), block_buffers.end(), [](const BlockBuffer &block) {
298+
return !IsChecksumHashableBlock(block);
299+
});
300+
if (invalid_it != block_buffers.end()) {
301+
const size_t idx = std::distance(block_buffers.begin(), invalid_it);
302+
KVCM_LOG_WARN("block [%zu] has ignored, empty, null, or zero-size iovs; skip checksum print", idx);
303+
} else if (sdk_buffer_check_pool_) {
304+
auto handle = sdk_buffer_check_pool_->GetCell();
305+
std::vector<int64_t> block_checksums;
306+
if (HashBlocksByIovShape(block_buffers, handle, max_check_iov_num_, block_checksums)) {
307+
PrintBlockChecksumAndUri("get_", uri_str_vec, block_checksums, trace_info);
308+
} else {
309+
KVCM_LOG_WARN("checksum print failed to hash blocks safely; skip checksum print");
310+
}
311+
} else {
312+
KVCM_LOG_WARN("KVCM_SDK_CHECK is enabled but sdk_buffer_check_pool is not initialized; "
313+
"skip checksum print");
314+
}
301315
}
302-
PrintBlockChecksumAndUri("get_", uri_str_vec, block_checksums, trace_info);
303316
}
304317
// Read-side verification path: only kicks in when caller supplies expected checksums
305318
// (typically forwarded from CacheLocation.checksum returned by meta service).

kv_cache_manager/client/test/meta_client_test.cc

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -127,7 +127,7 @@ class MockStub : public Stub {
127127
const std::string write_session_id,
128128
const BlockMask &success_block,
129129
const Locations &locations,
130-
const std::vector<int64_t> &checksums),
130+
const FinishWriteOptions &options),
131131
(override));
132132

133133
MOCK_METHOD(ClientErrorCode,
@@ -489,7 +489,7 @@ TEST_F(MetaClientTest, TestFinishWriteOptionsForwardToStub) {
489489
write_session_id,
490490
success_block,
491491
locations,
492-
::testing::IsEmpty()))
492+
::testing::Field(&FinishWriteOptions::checksums, ::testing::IsEmpty())))
493493
.Times(1)
494494
.WillOnce(::testing::Return(ER_OK));
495495

@@ -502,7 +502,7 @@ TEST_F(MetaClientTest, TestFinishWriteOptionsForwardToStub) {
502502
write_session_id,
503503
success_block,
504504
locations,
505-
::testing::ElementsAre(0x44, 0)))
505+
::testing::Field(&FinishWriteOptions::checksums, ::testing::ElementsAre(0x44, 0))))
506506
.Times(1)
507507
.WillOnce(::testing::Return(ER_OK));
508508

kv_cache_manager/manager/cache_manager.cc

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -701,7 +701,11 @@ CacheManager::StartWriteCache(RequestContext *request_context,
701701
RequestContext temp_request_context(trace_id + "_timeout_callback");
702702
BlockMaskOffset succeed_block = 0;
703703
auto ec = this->FinishWriteCache(
704-
&temp_request_context, instance_id, write_session_id, succeed_block, std::move(write_location_info));
704+
&temp_request_context,
705+
instance_id,
706+
write_session_id,
707+
succeed_block,
708+
CacheManager::FinishWriteCacheOptions::WithWriteLocationInfo(std::move(write_location_info)));
705709
static_cast<void>(ec);
706710
});
707711
KVCM_METRICS_COLLECTOR_CHRONO_MARK_END(service_metrics_collector, PutWriteLocationManager);
@@ -723,14 +727,13 @@ CacheManager::FinishWriteCache(RequestContext *request_context,
723727
const std::string &instance_id,
724728
const std::string &write_session_id,
725729
const BlockMask &success_block_mask,
726-
std::unique_ptr<WriteLocationManager::WriteLocationInfo> write_location_info_internal,
727-
const std::vector<int64_t> &checksums) {
730+
CacheManager::FinishWriteCacheOptions options) {
728731
SPAN_TRACER(request_context);
729732
const std::string &trace_id = request_context->trace_id();
730733
auto *service_metrics_collector = dynamic_cast<ServiceMetricsCollector *>(request_context->metrics_collector());
731734
WriteLocationManager::WriteLocationInfo location_info;
732-
if (write_location_info_internal != nullptr) {
733-
location_info = std::move(*write_location_info_internal);
735+
if (options.write_location_info_internal != nullptr) {
736+
location_info = std::move(*options.write_location_info_internal);
734737
} else if (!write_location_manager_->GetAndDelete(write_session_id, location_info)) {
735738
request_context->error_tracer()->AddErrorMsg("write_session_id has been deleted");
736739
RETURN_IF_EC_NOT_OK_WITH_LOG(
@@ -746,12 +749,12 @@ CacheManager::FinishWriteCache(RequestContext *request_context,
746749
// checksums must be the same length as keys (full batch, not the success subset).
747750
// An empty vector means the client did not report any; existing CacheLocation
748751
// checksums are then preserved. Any length mismatch is treated as a client bug.
749-
const bool has_checksums = !checksums.empty();
750-
if (has_checksums && checksums.size() != location_info.keys.size()) {
752+
const bool has_checksums = !options.checksums.empty();
753+
if (has_checksums && options.checksums.size() != location_info.keys.size()) {
751754
RETURN_IF_EC_NOT_OK_WITH_LOG(WARN,
752755
EC_BADARGS,
753756
"checksums size (%zu) does not match keys size (%zu)",
754-
checksums.size(),
757+
options.checksums.size(),
755758
location_info.keys.size());
756759
}
757760

@@ -771,7 +774,7 @@ CacheManager::FinishWriteCache(RequestContext *request_context,
771774
MetaSearcher::LocationUpdateTask task{
772775
.location_id = location_info.location_ids[block_key_idx],
773776
.new_status = CacheLocationStatus::CLS_SERVING,
774-
.checksum = has_checksums ? checksums[block_key_idx] : int64_t{0},
777+
.checksum = has_checksums ? options.checksums[block_key_idx] : int64_t{0},
775778
};
776779
success_batch_update_tasks.push_back({std::move(task)});
777780
} else {

kv_cache_manager/manager/cache_manager.h

Lines changed: 24 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
#include <memory>
66
#include <string>
77
#include <thread>
8+
#include <utility>
89
#include <vector>
910

1011
#include "kv_cache_manager/common/error_code.h"
@@ -144,17 +145,34 @@ class CacheManager {
144145
const std::vector<std::string> &location_spec_group_names,
145146
int64_t write_timeout_seconds,
146147
int32_t min_replica_count = 1);
147-
// checksums must be the same length as the keys captured at StartWriteCache (full
148-
// batch, not just the success subset within success_block_mask). Pass an empty
149-
// vector when the client does not report a checksum; existing CacheLocation
150-
// checksums are then preserved.
148+
struct FinishWriteCacheOptions {
149+
std::unique_ptr<WriteLocationManager::WriteLocationInfo> write_location_info_internal;
150+
// checksums must be the same length as the keys captured at StartWriteCache
151+
// (full batch, not just the success subset within success_block_mask). Pass
152+
// an empty vector when the client does not report a checksum; existing
153+
// CacheLocation checksums are then preserved.
154+
std::vector<int64_t> checksums;
155+
156+
static FinishWriteCacheOptions WithWriteLocationInfo(
157+
std::unique_ptr<WriteLocationManager::WriteLocationInfo> write_location_info_internal) {
158+
FinishWriteCacheOptions options;
159+
options.write_location_info_internal = std::move(write_location_info_internal);
160+
return options;
161+
}
162+
163+
static FinishWriteCacheOptions WithChecksums(std::vector<int64_t> checksums) {
164+
FinishWriteCacheOptions options;
165+
options.checksums = std::move(checksums);
166+
return options;
167+
}
168+
};
169+
151170
ErrorCode
152171
FinishWriteCache(RequestContext *request_context,
153172
const std::string &instance_id,
154173
const std::string &write_session_id,
155174
const BlockMask &success_block_mask,
156-
std::unique_ptr<WriteLocationManager::WriteLocationInfo> write_location_info_internal = nullptr,
157-
const std::vector<int64_t> &checksums = {});
175+
FinishWriteCacheOptions options = FinishWriteCacheOptions{});
158176

159177
ErrorCode RemoveCache(RequestContext *request_context,
160178
const std::string &instance_id,

0 commit comments

Comments
 (0)