From 3cbf58f68e483b3ea99bd3bd31705a72dc7c599b Mon Sep 17 00:00:00 2001 From: Googler Date: Thu, 17 Sep 2026 19:18:24 -0700 Subject: [PATCH] Add TdsKVBackend storage driver and wire TDS backend into KVCacheManagerBase and KVCacheStore. Standard buffered POSIX file I/O copies KV cache blocks through the Linux kernel page cache, incurring CPU `memcpy` overhead and host RAM page-cache bloat on high-throughput storage tiers such as Lustre, NFS, and NVMe SSDs. Implement `tpu_raiden::kv_cache::backends::storage::TdsKVBackend` (`tkv::backends::storage::TdsKVBackend`) inheriting from `KVBackend`, and wire `"tds"` backend creation and Host DRAM pool buffer registration (`RegisterBuffer`) into `KVCacheManagerBase` and `KVCacheStore`. All registration and data movement is delegated to `libtdsul` rather than to a hand-rolled I/O engine. Host buffers are registered with `tds_buffer_register_vaddr` / `tds_buffer_register_dmabuf` using `TDS_MEM_HOST`, the target file is registered with `tds_storage_handle_register` through an `fd://` URI, and transfers are issued via `tds_read` / `tds_write` / `tds_readv` / `tds_writev`. `TdsKVBackend` opens the storage file with `O_DIRECT` and hands that descriptor to `libtdsul`. Because `O_DIRECT` is a property of the open file description, the underlying `pread` / `pwrite` inherit it and bypass the kernel page cache. For 4KB-aligned, registered slices this is a zero-bounce transfer (`direct_ops`). Unaligned buffer addresses, slice lengths, or file offsets are staged through a 4KB-aligned bounce buffer owned by `TdsKVBackend` (`bounced_ops`), which performs a read-modify-write to preserve surrounding bytes and an `ftruncate` to preserve exact logical file sizes. Per `go/tdsul-api`, bounce staging belongs in the framework adapter and never inside `libtdsul`. A separate `o_direct_ops` counter measures whether the page-cache bypass was actually retained, because both the open path and the I/O path can silently fall back to buffered mode. Vendor `libtdsul` under `third_party/tpu_raiden/tpu_sync/tpudirect_storage/` so that non-experimental `tpu_raiden` targets do not depend on an `//experimental/...` package. This copy is an interim measure; see the TODO in its BUILD file. It also carries three fixes to `ThreadPoolConductor` that upstream does not have: worker CPU affinity is applied only when an explicit core range is configured, `pthread_create` failures retry unpinned instead of leaving a silently dead worker, and `Stop()` drains the pending task queue instead of leaking queued requests. PiperOrigin-RevId: 983561034 --- tpu_sync/core/controller/BUILD | 3 + .../core/controller/raiden_controller_test.cc | 68 +- tpu_sync/core/controller/test_util.h | 24 +- tpu_sync/kv_cache/BUILD | 5 + tpu_sync/kv_cache/backends/storage/BUILD | 87 ++ .../backends/storage/storage_backend_utils.cc | 56 ++ .../backends/storage/storage_backend_utils.h | 36 + .../kv_cache/backends/storage/tds_backend.cc | 666 ++++++++++++++ .../kv_cache/backends/storage/tds_backend.h | 207 +++++ .../backends/storage/tds_backend_test.cc | 747 ++++++++++++++++ tpu_sync/kv_cache/kv_cache_manager_base.cc | 34 +- tpu_sync/kv_cache/kv_cache_store_test.cc | 199 ++++- tpu_sync/tpudirect_storage/BUILD | 190 ++++ .../tpudirect_storage/include/tdsul/def.h | 155 ++++ .../tpudirect_storage/include/tdsul/tdsul.h | 143 +++ tpu_sync/tpudirect_storage/src/def_internal.h | 123 +++ .../tpudirect_storage/src/syscall_internal.h | 49 + tpu_sync/tpudirect_storage/src/tdsul.cpp | 841 ++++++++++++++++++ .../tpudirect_storage/src/util/IoConductor.h | 35 + .../tpudirect_storage/src/util/IoQueue.cpp | 98 ++ tpu_sync/tpudirect_storage/src/util/IoQueue.h | 98 ++ .../src/util/IoUringConductor.cpp | 44 + .../src/util/IoUringConductor.h | 39 + .../src/util/ThreadPoolConductor.cpp | 240 +++++ .../src/util/ThreadPoolConductor.h | 91 ++ .../tpudirect_storage/tests/io_queue_test.cpp | 244 +++++ .../tpudirect_storage/tests/tdsul_test.cpp | 492 ++++++++++ .../tests/thread_pool_conductor_test.cpp | 125 +++ 28 files changed, 5103 insertions(+), 36 deletions(-) create mode 100644 tpu_sync/kv_cache/backends/storage/storage_backend_utils.cc create mode 100644 tpu_sync/kv_cache/backends/storage/storage_backend_utils.h create mode 100644 tpu_sync/kv_cache/backends/storage/tds_backend.cc create mode 100644 tpu_sync/kv_cache/backends/storage/tds_backend.h create mode 100644 tpu_sync/kv_cache/backends/storage/tds_backend_test.cc create mode 100644 tpu_sync/tpudirect_storage/BUILD create mode 100644 tpu_sync/tpudirect_storage/include/tdsul/def.h create mode 100644 tpu_sync/tpudirect_storage/include/tdsul/tdsul.h create mode 100644 tpu_sync/tpudirect_storage/src/def_internal.h create mode 100644 tpu_sync/tpudirect_storage/src/syscall_internal.h create mode 100644 tpu_sync/tpudirect_storage/src/tdsul.cpp create mode 100644 tpu_sync/tpudirect_storage/src/util/IoConductor.h create mode 100644 tpu_sync/tpudirect_storage/src/util/IoQueue.cpp create mode 100644 tpu_sync/tpudirect_storage/src/util/IoQueue.h create mode 100644 tpu_sync/tpudirect_storage/src/util/IoUringConductor.cpp create mode 100644 tpu_sync/tpudirect_storage/src/util/IoUringConductor.h create mode 100644 tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.cpp create mode 100644 tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.h create mode 100644 tpu_sync/tpudirect_storage/tests/io_queue_test.cpp create mode 100644 tpu_sync/tpudirect_storage/tests/tdsul_test.cpp create mode 100644 tpu_sync/tpudirect_storage/tests/thread_pool_conductor_test.cpp diff --git a/tpu_sync/core/controller/BUILD b/tpu_sync/core/controller/BUILD index 70779228f..b3d056341 100644 --- a/tpu_sync/core/controller/BUILD +++ b/tpu_sync/core/controller/BUILD @@ -48,6 +48,7 @@ cc_library( "//tpu_sync/core:transfer_program_reshard", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/proto:transfer_program_cc_proto", "//tpu_sync/proto:worker_service_cc_grpc", "//tpu_sync/proto:worker_service_cc_proto", @@ -217,6 +218,7 @@ cc_library( "//tpu_sync/kv_cache:kv_cache_store_backend_factory", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "@com_github_grpc_grpc//:grpc++", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log:check", @@ -262,6 +264,7 @@ cc_test( "//tpu_sync/kv_cache:kv_cache_store_backend_factory", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/proto:worker_service_cc_proto", "//tpu_sync/rpc:raiden_service_cc_proto", "@com_google_absl//absl/container:flat_hash_map", diff --git a/tpu_sync/core/controller/raiden_controller_test.cc b/tpu_sync/core/controller/raiden_controller_test.cc index 1688b2aba..883c09517 100644 --- a/tpu_sync/core/controller/raiden_controller_test.cc +++ b/tpu_sync/core/controller/raiden_controller_test.cc @@ -20,7 +20,6 @@ #include #include #include -#include #include #include #include @@ -1516,6 +1515,73 @@ TEST_F(RaidenControllerTest, TransferBackendBuffersDispatchesOffloadAndRecall) { EXPECT_THAT(mock_mgr.last_backend_dst_block_ids, ElementsAre(4)); } +// ----------------------------------------------------------------------------- +// Verifies that `RaidenController::TransferBackendBuffers` dispatches both +// OFFLOAD ([hbm] -> [dram: User host_buf] -> storage) and RECALL +// (storage -> [dram: User host_buf] -> [hbm]) to the worker when the secondary +// backend is `"tds"` (`TdsKVBackend`). +// +// Specifically checks that: +// 1. `MockTransferManager::RegisterKVBackends` (test_util.h) creates a +// `TdsKVBackend` for backend type `"tds"`. +// 2. For `TRANSFER_DIR_OFFLOAD`, the worker maps `"tds_block0"` through the +// backend's `mapper()->MapKey()` to a path ending in the hex encoding +// (`7464735f626c6f636b30.bin`) and calls `D2hWriteToBackend` once. +// 3. For `TRANSFER_DIR_RECALL`, the worker maps `"tds_recall0"` to +// `7464735f726563616c6c30.bin` and calls `H2dReadFromBackend` once. +// +// The mock only records the calls; no storage I/O is performed. +// ----------------------------------------------------------------------------- +TEST_F(RaidenControllerTest, + TransferBackendBuffersDispatchesOffloadAndRecallWithTdsBackend) { + // Step 1: Register the `"tds"` secondary storage backend on the worker's + // transfer manager so `GetKVBackend("tds")` returns a `TdsKVBackend`. + kv_cache::BackendConfig tds_cfg; + tds_cfg.type = "tds"; + tds_cfg.parallelism.tp_rank = 0; + tds_cfg.parallelism.tp_size = 1; + tds_cfg.SetProperty("tp_size", "1"); + + MockTransferManager mock_mgr; + mock_mgr.RegisterKVBackends({tds_cfg}); + test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr)); + + TF_ASSERT_OK_AND_ASSIGN( + auto controller, + RaidenController::Create(unit_, /*num_blocks=*/5, /*num_shards=*/1, + /*shard_size_bytes=*/512, "")); + RegisterAndInitWorker(*controller, "worker_0", test_server_->server_address); + + ::tpu_sync::proto::BackendTransferSpec backend_spec; + backend_spec.set_name("tds"); + + // Step 2: Dispatch TRANSFER_DIR_OFFLOAD (`[hbm]` block 2 -> `[dram: User + // host_buf]` staging block 3 -> `"tds"` storage backend) and verify the + // resolved hex storage path passed to `D2HAndWriteToBackend`. + auto offload_status = controller->TransferBackendBuffers( + ::tpu_sync::proto::TRANSFER_DIR_OFFLOAD, {"tds_block0"}, + /*hbm_block_ids=*/{2}, /*host_block_ids=*/{3}, {backend_spec}); + ABSL_EXPECT_OK(offload_status.Await()); + EXPECT_EQ(mock_mgr.d2h_write_to_backend_calls, 1); + ASSERT_EQ(mock_mgr.last_d2h_backend_keys.size(), 1); + EXPECT_EQ(mock_mgr.last_d2h_backend_keys[0].block_hash, "tds_block0"); + EXPECT_THAT(mock_mgr.last_d2h_backend_keys[0].resolved_key, + HasSubstr("7464735f626c6f636b30.bin")); // hex("tds_block0") + + // Step 3: Dispatch TRANSFER_DIR_RECALL (`"tds"` storage backend -> `[dram: + // User host_buf]` staging block 1 -> `[hbm]` block 4) and verify the + // resolved hex storage path passed to `ReadFromBackendAndH2D`. + auto recall_status = controller->TransferBackendBuffers( + ::tpu_sync::proto::TRANSFER_DIR_RECALL, {"tds_recall0"}, + /*hbm_block_ids=*/{4}, /*host_block_ids=*/{1}, {backend_spec}); + ABSL_EXPECT_OK(recall_status.Await()); + EXPECT_EQ(mock_mgr.h2d_read_from_backend_calls, 1); + ASSERT_EQ(mock_mgr.last_h2d_backend_keys.size(), 1); + EXPECT_EQ(mock_mgr.last_h2d_backend_keys[0].block_hash, "tds_recall0"); + EXPECT_THAT(mock_mgr.last_h2d_backend_keys[0].resolved_key, + HasSubstr("7464735f726563616c6c30.bin")); // hex("tds_recall0") +} + TEST_F(RaidenControllerTest, TransferBuffersBackendSpecValidationRejections) { MockTransferManager mock_mgr; test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr)); diff --git a/tpu_sync/core/controller/test_util.h b/tpu_sync/core/controller/test_util.h index fc7769d8d..e1c98bd3b 100644 --- a/tpu_sync/core/controller/test_util.h +++ b/tpu_sync/core/controller/test_util.h @@ -44,6 +44,7 @@ #include "tpu_sync/core/raw_transfer_core.h" #include "tpu_sync/kv_cache/backends/backend.h" #include "tpu_sync/kv_cache/backends/storage/posix_backend.h" +#include "tpu_sync/kv_cache/backends/storage/tds_backend.h" #include "tpu_sync/kv_cache/kv_cache_store_backend_factory.h" namespace tpu_raiden { @@ -214,23 +215,32 @@ struct MockTransferManager { void RegisterKVBackends( absl::Span backend_configs) { for (const auto& cfg : backend_configs) { - if (!absl::EqualsIgnoreCase( - cfg.type, kv_cache::backends::storage::kPosixBackendName)) { + const bool is_posix = absl::EqualsIgnoreCase( + cfg.type, kv_cache::backends::storage::kPosixBackendName); + const bool is_tds = absl::EqualsIgnoreCase( + cfg.type, kv_cache::backends::storage::kTdsBackendName); + if (!is_posix && !is_tds) { continue; } if (cfg.parallelism.tp_rank < 0) continue; const std::string canonical_name = - std::string(kv_cache::backends::storage::kPosixBackendName); + std::string(is_tds ? kv_cache::backends::storage::kTdsBackendName + : kv_cache::backends::storage::kPosixBackendName); if (GetKVBackend(canonical_name) != nullptr) continue; kv_cache::BackendConfig resolved = cfg; kv_cache::ApplyParallelismToProperties( {cfg.parallelism.tp_size > 0 ? cfg.parallelism.tp_size : 1, cfg.parallelism.tp_rank}, &resolved); - auto backend = - std::make_shared( - canonical_name, resolved.properties); - backends[canonical_name] = std::move(backend); + if (is_tds) { + backends[canonical_name] = + std::make_shared( + canonical_name, resolved.properties); + } else { + backends[canonical_name] = + std::make_shared( + canonical_name, resolved.properties); + } } } diff --git a/tpu_sync/kv_cache/BUILD b/tpu_sync/kv_cache/BUILD index e3c1a9e7c..02b404a5f 100644 --- a/tpu_sync/kv_cache/BUILD +++ b/tpu_sync/kv_cache/BUILD @@ -149,6 +149,7 @@ cc_library( ":kv_cache_store_backend_factory", "//tpu_sync/common:raiden_id", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "@com_google_absl//absl/log", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", @@ -170,6 +171,7 @@ cc_test( ":kv_cache_store_wrapper", "//tpu_sync/common:raiden_id", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/kv_cache/global_registry:test_util", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/strings", @@ -283,6 +285,7 @@ cc_library( "//tpu_sync/core:xla_raw_transfer_headers", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/rpc:raiden_service_cc_proto", "//tpu_sync/telemetry:metrics_api", "//tpu_sync/telemetry:metrics_backend", @@ -562,6 +565,7 @@ cc_library( "//tpu_sync/core/controller:raiden_controller", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/kv_cache/global_registry:global_registry_client_cc", "//tpu_sync/kv_cache/reshard:reshard_service", "//tpu_sync/rpc:raiden_service_cc_proto", @@ -611,6 +615,7 @@ cc_test( "//tpu_sync/core/controller:test_util", "//tpu_sync/kv_cache/backends:backend", "//tpu_sync/kv_cache/backends/storage:posix_backend", + "//tpu_sync/kv_cache/backends/storage:tds_backend", "//tpu_sync/kv_cache/global_registry:global_registry_cc_grpc", "//tpu_sync/kv_cache/global_registry:global_registry_client_cc", "//tpu_sync/kv_cache/global_registry:global_registry_server_lib", diff --git a/tpu_sync/kv_cache/backends/storage/BUILD b/tpu_sync/kv_cache/backends/storage/BUILD index 4e6907fda..e2eea5347 100644 --- a/tpu_sync/kv_cache/backends/storage/BUILD +++ b/tpu_sync/kv_cache/backends/storage/BUILD @@ -87,3 +87,90 @@ cc_test( "@xla//xla/tsl/platform:statusor", ], ) + +cc_library( + name = "storage_backend_utils", + srcs = ["storage_backend_utils.cc"], + hdrs = ["storage_backend_utils.h"], + copts = [ + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + "-layering_check", + ], + visibility = ["//visibility:public"], + deps = [ + "//tpu_sync/kv_cache/backends:backend", + "@com_google_absl//absl/strings", + "@com_google_absl//absl/strings:string_view", + "@com_google_absl//absl/types:span", + ], +) + +cc_library( + name = "tds_backend", + srcs = ["tds_backend.cc"], + hdrs = ["tds_backend.h"], + copts = [ + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + "-layering_check", + ], + visibility = ["//visibility:public"], + deps = [ + ":storage_backend_utils", + "//tpu_sync/common:raiden_id", + "//tpu_sync/core:numa_thread_pool", + "//tpu_sync/core/controller:raiden_controller", + "//tpu_sync/kv_cache:kv_cache_store_backend", + "//tpu_sync/kv_cache:kv_cache_store_backend_factory", + "//tpu_sync/kv_cache/backends:backend", + "//tpu_sync/tpudirect_storage:tdsul", + "@com_google_absl//absl/base", + "@com_google_absl//absl/base:nullability", + "@com_google_absl//absl/cleanup", + "@com_google_absl//absl/container:flat_hash_map", + "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_macros", + "@com_google_absl//absl/status:statusor", + "@com_google_absl//absl/strings", + "@com_google_absl//absl/strings:string_view", + "@com_google_absl//absl/types:span", + "@xla//xla/tsl/concurrency:future", + "@xla//xla/tsl/platform:logging", + ], + alwayslink = 1, +) + +cc_test( + name = "tds_backend_test", + srcs = ["tds_backend_test.cc"], + copts = [ + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + "-layering_check", + ], + deps = [ + ":storage_backend_utils", + ":tds_backend", + "//tpu_sync/kv_cache:kv_cache_store_backend", + "//tpu_sync/kv_cache:kv_cache_store_backend_factory", + "//tpu_sync/kv_cache/backends:backend", + "@com_google_absl//absl/container:flat_hash_map", + "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_matchers", + "@com_google_absl//absl/status:statusor", + "@com_google_absl//absl/strings", + "@com_google_absl//absl/synchronization", + "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest_main", + ], +) diff --git a/tpu_sync/kv_cache/backends/storage/storage_backend_utils.cc b/tpu_sync/kv_cache/backends/storage/storage_backend_utils.cc new file mode 100644 index 000000000..d37421648 --- /dev/null +++ b/tpu_sync/kv_cache/backends/storage/storage_backend_utils.cc @@ -0,0 +1,56 @@ +// Copyright 2026 Google LLC. +// +// 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 "tpu_sync/kv_cache/backends/storage/storage_backend_utils.h" + +#include + +#include +#include +#include + +#include "absl/strings/ascii.h" +#include "absl/strings/string_view.h" +#include "absl/types/span.h" +#include "tpu_sync/kv_cache/backends/backend.h" + +namespace tpu_raiden::kv_cache::backends::storage { + +size_t GetStorageDirectIOAlignment() { + static const size_t kAlign = static_cast(sysconf(_SC_PAGESIZE)); + return kAlign; +} + +bool AreStorageSlicesDirectIOAligned( + absl::Span slices) { + const uintptr_t mask = GetStorageDirectIOAlignment() - 1; + for (const auto& s : slices) { + if (((reinterpret_cast(s.ptr) | s.size) & mask) != 0) + return false; + } + return true; +} + +std::string SanitizeStorageModelName(absl::string_view model_name) { + std::string out; + for (const char c : model_name) { + out.push_back((absl::ascii_isalnum(static_cast(c)) || + c == '.' || c == '_' || c == '-') + ? c + : '_'); + } + return out.empty() ? "unknown" : out; +} + +} // namespace tpu_raiden::kv_cache::backends::storage diff --git a/tpu_sync/kv_cache/backends/storage/storage_backend_utils.h b/tpu_sync/kv_cache/backends/storage/storage_backend_utils.h new file mode 100644 index 000000000..8845e8b2e --- /dev/null +++ b/tpu_sync/kv_cache/backends/storage/storage_backend_utils.h @@ -0,0 +1,36 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_STORAGE_BACKEND_UTILS_H_ +#define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_STORAGE_BACKEND_UTILS_H_ + +#include +#include + +#include "absl/strings/string_view.h" +#include "absl/types/span.h" +#include "tpu_sync/kv_cache/backends/backend.h" + +namespace tpu_raiden::kv_cache::backends::storage { + +size_t GetStorageDirectIOAlignment(); + +bool AreStorageSlicesDirectIOAligned( + absl::Span slices); + +std::string SanitizeStorageModelName(absl::string_view model_name); + +} // namespace tpu_raiden::kv_cache::backends::storage + +#endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_STORAGE_BACKEND_UTILS_H_ diff --git a/tpu_sync/kv_cache/backends/storage/tds_backend.cc b/tpu_sync/kv_cache/backends/storage/tds_backend.cc new file mode 100644 index 000000000..c6b488769 --- /dev/null +++ b/tpu_sync/kv_cache/backends/storage/tds_backend.cc @@ -0,0 +1,666 @@ +// Copyright 2026 Google LLC. +// +// 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 "tpu_sync/kv_cache/backends/storage/tds_backend.h" + +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include // NOLINT(build/c++17) +#include +#include +#include +#include +#include +#include // NOLINT(build/c++11) +#include +#include + +#include "tdsul/def.h" +#include "tdsul/tdsul.h" +#include "absl/base/call_once.h" +#include "absl/container/flat_hash_map.h" +#include "absl/status/status.h" +#include "absl/status/status_macros.h" +#include "absl/status/statusor.h" +#include "absl/strings/ascii.h" +#include "absl/strings/escaping.h" +#include "absl/strings/numbers.h" +#include "absl/strings/str_cat.h" +#include "absl/strings/string_view.h" +#include "absl/types/span.h" +#include "xla/tsl/platform/logging.h" +#include "tpu_sync/common/raiden_id.h" +#include "tpu_sync/core/controller/raiden_controller.h" +#include "tpu_sync/core/numa_thread_pool.h" +#include "tpu_sync/kv_cache/backends/backend.h" +#include "tpu_sync/kv_cache/backends/storage/storage_backend_utils.h" +#include "tpu_sync/kv_cache/kv_cache_store_backend.h" +#include "tpu_sync/kv_cache/kv_cache_store_backend_factory.h" + +namespace tpu_raiden { +namespace kv_cache { +namespace backends { +namespace storage { + +namespace fs = std::filesystem; + +namespace { + +#ifndef UIO_MAXIOV +#define UIO_MAXIOV 1024 +#endif + +constexpr size_t kTdsDefaultLookupBatchSize = 32; +constexpr size_t kTdsHashDirL1Width = 3; +constexpr size_t kTdsHashDirL2Width = 2; +constexpr size_t kTdsMaxBlockHashBytes = 125; + +// Initializes the process-wide `libtdsul` runtime (`tds_init`) once. +void EnsureTdsulInitialized(int num_worker_threads) { + static absl::once_flag init_once; + absl::call_once(init_once, [num_worker_threads]() { + std::unique_ptr cfg( + tds_config_create(), &tds_config_destroy); + if (cfg == nullptr) { + LOG(ERROR) << "tds_config_create failed; libtdsul left uninitialized"; + return; + } + tds_config_set_int(cfg.get(), "num_worker_threads", + std::max(1, num_worker_threads)); + tds_config_set_int(cfg.get(), "enable_io_uring", 0); + tds_config_set_int(cfg.get(), "enable_p2p", 0); + const int rc = tds_init(cfg.get()); + if (rc != TDS_SUCCESS) { + LOG(ERROR) << "tds_init failed (" << rc << ")"; + } + }); +} + +size_t GetBufferIoSize(const tds_buffer_io_t& bio) { + return (bio.handle != nullptr) ? bio.registered.size : bio.raw_ptr.size; +} + +void AdvanceBufferIo(tds_buffer_io_t* bio, size_t consumed) { + if (bio->handle != nullptr) { + bio->registered.offset += consumed; + bio->registered.size -= consumed; + } else { + bio->raw_ptr.vptr = static_cast(bio->raw_ptr.vptr) + consumed; + bio->raw_ptr.size -= consumed; + } +} + +absl::Status ExecuteTdsIo(int fd, std::vector buffer_ios, + int64_t file_offset, bool is_write, + size_t min_required_bytes) { + if (buffer_ios.empty() && min_required_bytes == 0) { + return absl::OkStatus(); + } + + std::string fd_uri = absl::StrCat("fd://", fd); + tds_storage_descr_t storage_descr = { + .uri = fd_uri.c_str(), + .options = nullptr, + }; + tds_storage_handle_t* raw_storage_handle = nullptr; + if (tds_storage_handle_register(&storage_descr, &raw_storage_handle) != + TDS_SUCCESS || + raw_storage_handle == nullptr) { + return absl::InternalError( + absl::StrCat("tds_storage_handle_register failed for URI: ", fd_uri)); + } + std::unique_ptr + storage_handle(raw_storage_handle, &tds_storage_handle_deregister); + + size_t idx = 0; + size_t total_transferred = 0; + int64_t current_offset = file_offset; + + while (idx < buffer_ios.size()) { + if (GetBufferIoSize(buffer_ios[idx]) == 0) { + ++idx; + continue; + } + const int iov_cnt = std::min(buffer_ios.size() - idx, UIO_MAXIOV); + size_t batch_bytes = 0; + for (int i = 0; i < iov_cnt; ++i) { + batch_bytes += GetBufferIoSize(buffer_ios[idx + i]); + } + + tds_storage_io_t storage_io = + tds_create_file_io(storage_handle.get(), current_offset, batch_bytes); + + ssize_t n = -1; + if (iov_cnt == 1) { + n = is_write ? tds_write(&storage_io, &buffer_ios[idx]) + : tds_read(&storage_io, &buffer_ios[idx]); + } else { + n = is_write ? tds_writev(&storage_io, &buffer_ios[idx], iov_cnt) + : tds_readv(&storage_io, &buffer_ios[idx], iov_cnt); + } + + if (n < 0) { + if (errno == EINTR) continue; + return absl::ErrnoToStatus( + errno, is_write ? "libtdsul tds_write/tds_writev failed" + : "libtdsul tds_read/tds_readv failed"); + } + if (n == 0) { + break; + } + + size_t step = static_cast(n); + total_transferred += step; + current_offset += n; + while (step > 0 && idx < buffer_ios.size()) { + const size_t cur_len = GetBufferIoSize(buffer_ios[idx]); + if (step >= cur_len) { + step -= cur_len; + ++idx; + } else { + AdvanceBufferIo(&buffer_ios[idx], step); + step = 0; + } + } + } + + if (total_transferred < min_required_bytes) { + return absl::OutOfRangeError(absl::StrCat( + "Short transfer in libtdsul I/O: required ", min_required_bytes, + " bytes, transferred ", total_transferred)); + } + return absl::OkStatus(); +} + +} // namespace + +// --- TdsBackendOptions Implementation --- + +absl::StatusOr TdsBackendOptions::FromProperties( + const absl::flat_hash_map& properties) { + TdsBackendOptions options; + auto str = [&](absl::string_view key, std::string* out) { + auto it = properties.find(key); + if (it != properties.end()) *out = it->second; + }; + auto num = [&](absl::string_view key, int64_t* out) -> absl::Status { + auto it = properties.find(key); + if (it == properties.end()) return absl::OkStatus(); + if (!absl::SimpleAtoi(it->second, out)) { + return absl::InvalidArgumentError( + absl::StrCat(key, " is not an integer: ", it->second)); + } + return absl::OkStatus(); + }; + str("root_dir", &options.root_dir); + str("model_name", &options.model_name); + int64_t tp_size = options.tp_size; + int64_t tp_rank = options.tp_rank; + int64_t capacity = 0; + int64_t batch = options.lookup_batch_size; + int64_t threads = options.storage_io_thread_pool_size; + ABSL_RETURN_IF_ERROR(num("tp_size", &tp_size)); + ABSL_RETURN_IF_ERROR(num("tp_rank", &tp_rank)); + ABSL_RETURN_IF_ERROR(num("capacity_bytes", &capacity)); + ABSL_RETURN_IF_ERROR(num("lookup_batch_size", &batch)); + ABSL_RETURN_IF_ERROR(num("storage_io_thread_pool_size", &threads)); + if (tp_size < 1) { + return absl::InvalidArgumentError( + absl::StrCat("tp_size must be >= 1, got ", tp_size)); + } + if (tp_rank < 0 || tp_rank >= tp_size) { + return absl::InvalidArgumentError( + absl::StrCat("tp_rank must be in [0, ", tp_size, "), got ", tp_rank)); + } + if (threads < 1) { + return absl::InvalidArgumentError(absl::StrCat( + "storage_io_thread_pool_size must be >= 1, got ", threads)); + } + options.tp_size = static_cast(tp_size); + options.tp_rank = static_cast(tp_rank); + options.capacity_bytes = static_cast(capacity); + options.lookup_batch_size = + batch > 0 ? static_cast(batch) : kTdsDefaultLookupBatchSize; + options.storage_io_thread_pool_size = static_cast(threads); + + if (auto it = properties.find("direct_io"); it != properties.end()) { + std::string val = absl::AsciiStrToLower(it->second); + if (val == "true") { + options.direct_io = true; + } else if (val == "false") { + options.direct_io = false; + } else { + return absl::InvalidArgumentError( + absl::StrCat("Invalid boolean value for direct_io: '", it->second, + "'; expected 'true' or 'false'")); + } + } + return options; +} + +// --- TdsKVBackend Implementation --- + +bool TdsKVBackend::ProbeDirectIO(absl::string_view dir) { + static std::atomic probe_seq{0}; + std::string probe_path = + absl::StrCat(dir, "/.tds_o_direct_probe_", getpid(), "_", + probe_seq.fetch_add(1, std::memory_order_relaxed)); + std::error_code ec; + fs::create_directories(std::string(dir), ec); + + int fd = + open(probe_path.c_str(), O_WRONLY | O_CREAT | O_TRUNC | O_DIRECT, 0644); + if (fd < 0) { + return false; + } + + const size_t align = GetStorageDirectIOAlignment(); + void* page = nullptr; + if (posix_memalign(&page, align, align) != 0) { + close(fd); + unlink(probe_path.c_str()); + return false; + } + std::memset(page, 0, align); + ssize_t written = write(fd, page, align); + free(page); + close(fd); + unlink(probe_path.c_str()); + return (written == static_cast(align)); +} + +TdsKVBackend::TdsKVBackend( + std::string name, absl::flat_hash_map properties) + : KVBackend(std::move(properties)), name_(std::move(name)) { + if (!properties_.contains("tp_rank")) { + properties_["tp_rank"] = "0"; + } + absl::StatusOr options = + TdsBackendOptions::FromProperties(properties_); + if (!options.ok()) { + LOG(FATAL) << "[TdsKVBackend] invalid configuration for " << name_ << ": " + << options.status(); + } + options_ = *std::move(options); + EnsureTdsulInitialized(options_.storage_io_thread_pool_size); + if (options_.direct_io) { + direct_io_supported_ = ProbeDirectIO(options_.root_dir); + if (!direct_io_supported_) { + LOG(WARNING) << "[TdsKVBackend] " << name_ + << ": O_DIRECT not supported on '" << options_.root_dir + << "'; falling back to buffered I/O."; + } + } + mapper_ = std::make_shared( + options_.root_dir, options_.model_name, options_.tp_size, + options_.tp_rank); + thread_pool_ = + std::make_unique(options_.storage_io_thread_pool_size); +} + +TdsKVBackend::~TdsKVBackend() { thread_pool_.reset(); } + +absl::StatusOr> TdsKVBackend::BuildBufferIos( + absl::Span slices) const { + std::vector bios; + bios.reserve(slices.size()); + for (const auto& slice : slices) { + if (slice.size == 0) continue; + if (slice.ptr == nullptr) { + return absl::InvalidArgumentError( + "Null slice pointer with non-zero size"); + } + bios.push_back(tds_create_raw_buffer_io(slice.ptr, slice.size)); + } + return bios; +} + +void TdsKVBackend::WriteAsync(const BlockKey& key, + absl::Span slices, + size_t total_bytes, + std::function callback) { + thread_pool_->Schedule( + std::nullopt, + [this, key, + slices = std::vector(slices.begin(), slices.end()), + total_bytes, callback = std::move(callback)]() { + auto run = [&]() -> absl::Status { + if (key.offset != 0) { + return absl::InvalidArgumentError(absl::StrCat( + "TdsKVBackend::WriteAsync requires offset 0 (atomic whole-file " + "publish); got ", + key.offset)); + } + size_t sum_bytes = 0; + for (const auto& s : slices) sum_bytes += s.size; + if (sum_bytes != total_bytes) { + return absl::InvalidArgumentError(absl::StrCat( + "total_bytes (", total_bytes, + ") does not match sum of slice sizes (", sum_bytes, ")")); + } + ABSL_ASSIGN_OR_RETURN(std::vector bios, + BuildBufferIos(slices)); + + static std::atomic tmp_seq{0}; + const std::string tmp_path = + absl::StrCat(key.resolved_key, ".tmp_", ::getpid(), "_", + tmp_seq.fetch_add(1, std::memory_order_relaxed)); + const bool use_direct = + direct_io_supported_ && AreStorageSlicesDirectIOAligned(slices); + const int open_flags = + O_WRONLY | O_CREAT | O_TRUNC | (use_direct ? O_DIRECT : 0); + + int fd = ::open(tmp_path.c_str(), open_flags, 0644); + if (fd < 0 && errno == ENOENT) { + std::string dir_path(TdsPathMapper::GetParentDir(key.resolved_key)); + if (!dir_path.empty()) { + std::error_code ec; + fs::create_directories(dir_path, ec); + if (ec && !fs::exists(dir_path, ec)) { + return absl::InternalError( + absl::StrCat("Failed to create directory: ", dir_path, + ", error: ", ec.message())); + } + fd = ::open(tmp_path.c_str(), open_flags, 0644); + } + } + if (fd < 0) { + return absl::ErrnoToStatus( + errno, + absl::StrCat("Failed to open file for write: ", tmp_path)); + } + + absl::Status io_status = + ExecuteTdsIo(fd, std::move(bios), + /*file_offset=*/0, /*is_write=*/true, total_bytes); + if (!io_status.ok()) { + ::close(fd); + ::unlink(tmp_path.c_str()); + return io_status; + } + if (::close(fd) < 0) { + int saved_errno = errno; + ::unlink(tmp_path.c_str()); + return absl::ErrnoToStatus(saved_errno, + "Failed to close file after write"); + } + if (::rename(tmp_path.c_str(), key.resolved_key.c_str()) != 0) { + int saved_errno = errno; + ::unlink(tmp_path.c_str()); + return absl::ErrnoToStatus( + saved_errno, + absl::StrCat("Failed to publish ", key.resolved_key)); + } + return absl::OkStatus(); + }; + absl::Status s = run(); + if (callback) callback(std::move(s)); + }); +} + +void TdsKVBackend::ReadAsync(const BlockKey& key, + absl::Span slices, + size_t total_bytes, + std::function callback) { + thread_pool_->Schedule( + std::nullopt, + [this, key, + slices = std::vector(slices.begin(), slices.end()), + total_bytes, callback = std::move(callback)]() { + auto run = [&]() -> absl::Status { + if (key.offset < 0) { + return absl::InvalidArgumentError("Negative key offset"); + } + size_t sum_bytes = 0; + for (const auto& s : slices) sum_bytes += s.size; + if (sum_bytes != total_bytes) { + return absl::InvalidArgumentError(absl::StrCat( + "total_bytes (", total_bytes, + ") does not match sum of slice sizes (", sum_bytes, ")")); + } + ABSL_ASSIGN_OR_RETURN(std::vector bios, + BuildBufferIos(slices)); + + const bool use_direct = + direct_io_supported_ && + ((key.offset & (GetStorageDirectIOAlignment() - 1)) == 0) && + AreStorageSlicesDirectIOAligned(slices); + const int open_flags = O_RDONLY | (use_direct ? O_DIRECT : 0); + int fd = ::open(key.resolved_key.c_str(), open_flags); + if (fd < 0) { + if (errno == ENOENT) { + return absl::NotFoundError( + absl::StrCat("Block file not found: ", key.resolved_key)); + } + return absl::ErrnoToStatus( + errno, absl::StrCat("Failed to open file for read: ", + key.resolved_key)); + } + + absl::Status io_status = + ExecuteTdsIo(fd, std::move(bios), key.offset, + /*is_write=*/false, total_bytes); + if (!io_status.ok()) { + ::close(fd); + return io_status; + } + if (::close(fd) < 0) { + return absl::ErrnoToStatus(errno, + "Failed to close file after read"); + } + return absl::OkStatus(); + }; + absl::Status s = run(); + if (callback) callback(std::move(s)); + }); +} + +absl::StatusOr TdsKVBackend::Exists(const BlockKey& key) { + if (::faccessat(AT_FDCWD, key.resolved_key.c_str(), F_OK, AT_EACCESS) == 0) { + return true; + } + if (errno == ENOENT || errno == ENOTDIR) { + return false; + } + return absl::ErrnoToStatus( + errno, absl::StrCat("faccessat(F_OK) failed for: ", key.resolved_key)); +} + +void TdsKVBackend::BatchExistsAsync( + absl::Span keys, + std::function>)> callback) { + if (keys.empty()) { + thread_pool_->Schedule(std::nullopt, [callback = std::move(callback)]() { + if (callback) callback({}); + }); + return; + } + + const size_t num_keys = keys.size(); + struct AsyncBatchState { + std::vector> results; + std::atomic remaining; + std::function>)> callback; + }; + + auto state = std::make_shared(); + state->results.resize(num_keys); + state->remaining.store(num_keys, std::memory_order_relaxed); + state->callback = std::move(callback); + + for (size_t i = 0; i < num_keys; ++i) { + thread_pool_->Schedule(std::nullopt, [this, key = keys[i], i, state]() { + state->results[i] = Exists(key); + if (state->remaining.fetch_sub(1, std::memory_order_acq_rel) == 1 && + state->callback) { + state->callback(std::move(state->results)); + } + }); + } +} + +// --- TdsPathMapper Implementation --- + +TdsPathMapper::TdsPathMapper(absl::string_view root_dir, + absl::string_view model_name, int tp_size, + int tp_rank) + : root_dir_(root_dir), + model_name_(SanitizeStorageModelName(model_name)), + tp_size_(tp_size), + tp_rank_(tp_rank) {} + +absl::StatusOr TdsPathMapper::MapKey( + const std::string& block_hash, const KeyMappingOptions& options) const { + if (block_hash.empty()) { + return absl::InvalidArgumentError("block_hash must not be empty."); + } + if (block_hash.size() > kTdsMaxBlockHashBytes) { + return absl::InvalidArgumentError( + absl::StrCat("block_hash is ", block_hash.size(), " bytes; maximum is ", + kTdsMaxBlockHashBytes, " (NAME_MAX after hex encoding).")); + } + const int target_rank = (options.parallelism.tp_rank == -1) + ? tp_rank_ + : options.parallelism.tp_rank; + const int target_tp_size = (options.parallelism.tp_size == -1) + ? tp_size_ + : options.parallelism.tp_size; + + const std::string hash_hex = absl::BytesToHexString(block_hash); + const std::string padded = absl::StrCat( + hash_hex, std::string(kTdsHashDirL1Width + kTdsHashDirL2Width, '0')); + const absl::string_view l1(padded.data(), kTdsHashDirL1Width); + const absl::string_view l2(padded.data() + kTdsHashDirL1Width, + kTdsHashDirL2Width); + + std::string resolved_path = + absl::StrCat(root_dir_, "/", model_name_, "/tp", target_tp_size, "_r", + target_rank, "/", l1, "/", l2, "/", hash_hex, ".bin"); + return BlockKey{block_hash, resolved_path, /*offset=*/0, /*size=*/0}; +} + +// --- TdsKVCacheStoreBackend Implementation --- + +TdsKVCacheStoreBackend::TdsKVCacheStoreBackend( + std::shared_ptr storage_backend, std::string name, + size_t capacity_bytes, size_t lookup_batch_size) + : storage_backend_(std::move(storage_backend)), + name_(std::move(name)), + capacity_bytes_(capacity_bytes), + lookup_batch_size_(lookup_batch_size > 0 ? lookup_batch_size + : kTdsDefaultLookupBatchSize) {} + +absl::StatusOr> TdsKVCacheStoreBackend::MapShardKeys( + const std::string& block_hash) const { + const std::shared_ptr mapper = storage_backend_->mapper(); + const int tp_size = mapper->tp_size(); + if (tp_size < 1) { + return absl::InvalidArgumentError( + absl::StrCat("invalid mapper tp_size: ", tp_size)); + } + std::vector keys; + keys.reserve(tp_size); + for (int rank = 0; rank < tp_size; ++rank) { + const backends::KeyMappingOptions lookup_opts{ + .parallelism = {.tp_size = tp_size, .tp_rank = rank}, + }; + ABSL_ASSIGN_OR_RETURN(BlockKey key, + mapper->MapKey(block_hash, lookup_opts)); + keys.push_back(std::move(key)); + } + return keys; +} + +absl::StatusOr TdsKVCacheStoreBackend::Lookup( + absl::Span block_hashes, const LookupOptions& options) { + BlockSliceList results; + if (!storage_backend_ || !storage_backend_->mapper() || block_hashes.empty()) + return results; + + for (const std::string& hash : block_hashes) { + absl::StatusOr> shard_keys = MapShardKeys(hash); + if (!shard_keys.ok()) break; + + std::promise>> promise; + auto future = promise.get_future(); + storage_backend_->BatchExistsAsync( + *shard_keys, [&promise](std::vector> res) { + promise.set_value(std::move(res)); + }); + const std::vector> answers = future.get(); + if (answers.size() != shard_keys->size()) break; + bool all_exist = true; + for (const auto& ans : answers) { + if (!ans.ok() || !*ans) { + all_exist = false; + break; + } + } + if (!all_exist) break; + + RaidenBlockId block; + block.status = BlockStatus::SHARED_STORAGE; + block.raiden_id.job_replica_id = "shared"; + block.raiden_id.data_name = name_; + results.push_back(std::make_pair(hash, block)); + } + + return results; +} + +} // namespace storage +} // namespace backends +} // namespace kv_cache +} // namespace tpu_raiden + +using ::tpu_raiden::kv_cache::backends::storage::TdsBackendOptions; +using ::tpu_raiden::kv_cache::backends::storage::TdsKVBackend; +using ::tpu_raiden::kv_cache::backends::storage::TdsKVCacheStoreBackend; + +REGISTER_KV_CACHE_STORE_BACKEND( + ::tpu_raiden::kv_cache::backends::storage::kTdsBackendName, + [](const ::tpu_raiden::kv_cache::BackendConfig& config, + ::tpu_raiden::controller::RaidenController* /*controller*/) + -> absl::StatusOr< + std::shared_ptr<::tpu_raiden::kv_cache::KVCacheStoreBackend>> { + const int tp_size = + config.parallelism.tp_size > 0 ? config.parallelism.tp_size : 1; + + ::tpu_raiden::kv_cache::BackendConfig resolved = config; + ::tpu_raiden::kv_cache::ApplyParallelismToProperties( + {.tp_size = tp_size, .tp_rank = 0}, &resolved); + ABSL_ASSIGN_OR_RETURN( + const TdsBackendOptions options, + TdsBackendOptions::FromProperties(resolved.properties)); + + const std::string backend_name = std::string( + ::tpu_raiden::kv_cache::backends::storage::kTdsBackendName); + auto backend = + std::make_shared(backend_name, resolved.properties); + return std::make_shared( + std::move(backend), backend_name, options.capacity_bytes, + options.lookup_batch_size); + }); diff --git a/tpu_sync/kv_cache/backends/storage/tds_backend.h b/tpu_sync/kv_cache/backends/storage/tds_backend.h new file mode 100644 index 000000000..ef7e5f5e8 --- /dev/null +++ b/tpu_sync/kv_cache/backends/storage/tds_backend.h @@ -0,0 +1,207 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_TDS_BACKEND_H_ +#define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_TDS_BACKEND_H_ + +#include +#include +#include +#include +#include +#include +#include + +#include "tdsul/tdsul.h" +#include "absl/base/nullability.h" +#include "absl/container/flat_hash_map.h" +#include "absl/status/status.h" +#include "absl/status/statusor.h" +#include "absl/strings/string_view.h" +#include "absl/types/span.h" +#include "xla/tsl/concurrency/future.h" +#include "tpu_sync/common/raiden_id.h" +#include "tpu_sync/core/numa_thread_pool.h" +#include "tpu_sync/kv_cache/backends/backend.h" +#include "tpu_sync/kv_cache/backends/storage/storage_backend_utils.h" +#include "tpu_sync/kv_cache/kv_cache_store_backend.h" + +namespace tpu_raiden { +namespace kv_cache { + +class BlockTracker; + +namespace backends { +namespace storage { + +inline constexpr absl::string_view kTdsBackendName = "tds"; + +// Typed, validated view of a TDS backend's configuration. +struct TdsBackendOptions { + std::string root_dir = "/tmp/raiden_storage"; + std::string model_name = "unknown"; + int tp_size = 1; + int tp_rank = -1; // Required; -1 means the caller did not supply it. + size_t capacity_bytes = 0; + size_t lookup_batch_size = 32; + int storage_io_thread_pool_size = 16; + bool direct_io = true; + + static absl::StatusOr FromProperties( + const absl::flat_hash_map& properties); +}; + +// TdsKVBackend implements TPU Raiden's KVBackend interface on top of +// `libtdsul` (`//tpu_sync/tpudirect_storage:tdsul`). +class TdsKVBackend : public KVBackend { + public: + explicit TdsKVBackend( + std::string name, + absl::flat_hash_map properties = {}); + + ~TdsKVBackend() override; + + std::string name() const override { return name_; } + + const TdsBackendOptions& options() const { return options_; } + + bool is_direct_io_supported() const { return direct_io_supported_; } + + static bool ProbeDirectIO(absl::string_view dir); + + void WriteAsync(const BlockKey& key, + absl::Span slices, + size_t total_bytes, + std::function callback) override; + + void ReadAsync(const BlockKey& key, + absl::Span slices, + size_t total_bytes, + std::function callback) override; + + void BatchExistsAsync( + absl::Span keys, + std::function>)> callback) override; + + private: + absl::StatusOr> BuildBufferIos( + absl::Span slices) const; + absl::StatusOr Exists(const BlockKey& key); + + std::string name_; + TdsBackendOptions options_; + bool direct_io_supported_ = false; + // MUST be the last-declared member so in-flight tasks finish before any + // member they touch is destroyed. + std::unique_ptr thread_pool_; +}; + +// TdsPathMapper maps block hash identifiers and tensor-parallel rank to a +// rank-partitioned hierarchical directory layout over the hex-encoded hash: +// `//tp_r///.bin` +class TdsPathMapper : public BlockKeyMapper { + public: + static absl::string_view GetParentDir(absl::string_view path) { + size_t last_slash = path.find_last_of('/'); + if (last_slash == absl::string_view::npos) return ""; + return path.substr(0, last_slash); + } + + TdsPathMapper(absl::string_view root_dir, absl::string_view model_name, + int tp_size, int tp_rank); + + absl::StatusOr MapKey( + const std::string& block_hash, + const KeyMappingOptions& options = {}) const override; + int tp_size() const override { return tp_size_; } + + private: + std::string root_dir_; + std::string model_name_; + int tp_size_; + int tp_rank_; +}; + +// TdsKVCacheStoreBackend probes persistent storage on the coordinator. +class TdsKVCacheStoreBackend : public KVCacheStoreBackend { + public: + explicit TdsKVCacheStoreBackend( + std::shared_ptr storage_backend, + std::string name = std::string(kTdsBackendName), + size_t capacity_bytes = 0, size_t lookup_batch_size = 32); + + std::string name() const override { return name_; } + + size_t lookup_batch_size() const { return lookup_batch_size_; } + + absl::StatusOr Lookup( + absl::Span block_hashes, + const LookupOptions& options = {}) override; + + tsl::Future<> Load(const RaidenId& remote_id, + absl::Span block_hashes, + absl::Span device_block_ids, + absl::Span slices, + BlockTracker* absl_nonnull load_tracker) override { + return tsl::Future<>(absl::OkStatus()); + } + + std::pair Insert( + absl::Span block_hashes, + absl::Span slices, bool on_host) override { + return {true, {}}; + } + + bool InsertAndLock(absl::Span block_hashes, + absl::Span slices, + bool on_host) override { + return true; + } + + size_t ReleaseAndDelete(absl::Span block_hashes) override { + return 0; + } + void Delete(absl::Span block_hashes, + absl::Span slices) override {} + bool Pin(absl::Span block_hashes) override { return true; } + void Release(absl::Span block_hashes) override {} + int GetPinCount(const std::string& hash) const override { return 0; } + size_t GetCapacity() const override { return capacity_bytes_; } + size_t GetSize() const override { return 0; } + size_t GetAvailableSpace() const override { return capacity_bytes_; } + + std::shared_ptr storage_backend() const { + return storage_backend_; + } + + private: + absl::StatusOr> MapShardKeys( + const std::string& block_hash) const; + + std::shared_ptr storage_backend_; + std::string name_ = std::string(kTdsBackendName); + size_t capacity_bytes_ = 0; + size_t lookup_batch_size_ = 32; +}; + +} // namespace storage +} // namespace backends +} // namespace kv_cache +} // namespace tpu_raiden + +namespace tkv { +namespace backends = ::tpu_raiden::kv_cache::backends; +} // namespace tkv + +#endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_KV_CACHE_BACKENDS_STORAGE_TDS_BACKEND_H_ diff --git a/tpu_sync/kv_cache/backends/storage/tds_backend_test.cc b/tpu_sync/kv_cache/backends/storage/tds_backend_test.cc new file mode 100644 index 000000000..74eb8efd2 --- /dev/null +++ b/tpu_sync/kv_cache/backends/storage/tds_backend_test.cc @@ -0,0 +1,747 @@ +// Copyright 2026 Google LLC. +// +// 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 "tpu_sync/kv_cache/backends/storage/tds_backend.h" + +#include +#include +#include + +#include +#include +#include +#include +#include +#include // NOLINT(build/c++17) +#include +#include +#include +#include +#include + +#include +#include +#include "absl/container/flat_hash_map.h" +#include "absl/status/status.h" +#include "absl/status/status_matchers.h" +#include "absl/status/statusor.h" +#include "absl/strings/str_cat.h" +#include "absl/synchronization/notification.h" +#include "absl/types/span.h" +#include "tpu_sync/kv_cache/backends/backend.h" +#include "tpu_sync/kv_cache/backends/storage/storage_backend_utils.h" +#include "tpu_sync/kv_cache/kv_cache_store_backend.h" +#include "tpu_sync/kv_cache/kv_cache_store_backend_factory.h" + +namespace tpu_raiden { +namespace kv_cache { +namespace backends { +namespace storage { +namespace { + +using ::absl_testing::IsOkAndHolds; +using ::absl_testing::StatusIs; + +static_assert(std::is_base_of_v); + +class AlignedBuffer { + public: + explicit AlignedBuffer(size_t size, size_t alignment = GetStorageDirectIOAlignment()) + : size_(size) { + void* raw = nullptr; + if (posix_memalign(&raw, alignment, size) != 0) { + raw = nullptr; + } + ptr_ = static_cast(raw); + } + + ~AlignedBuffer() { std::free(ptr_); } + + AlignedBuffer(const AlignedBuffer&) = delete; + AlignedBuffer& operator=(const AlignedBuffer&) = delete; + + uint8_t* data() const { return ptr_; } + size_t size() const { return size_; } + + private: + uint8_t* ptr_ = nullptr; + size_t size_ = 0; +}; + +std::vector MakeDeterministicPayload(size_t size, uint8_t seed) { + std::vector payload(size); + for (size_t i = 0; i < size; ++i) { + payload[i] = + static_cast((i * 37 + (i >> 8) * 101 + seed * 13 + 7) & 0xFF); + } + return payload; +} + +int64_t FileSizeOrNegative(const std::string& path) { + struct stat st = {}; + if (::stat(path.c_str(), &st) != 0) return -1; + return st.st_size; +} + +class TdsKVBackendTest : public ::testing::Test { + protected: + void SetUp() override { + const ::testing::TestInfo* info = + ::testing::UnitTest::GetInstance()->current_test_info(); + const std::string test_name = info != nullptr ? info->name() : "test"; + test_dir_ = absl::StrCat(::testing::TempDir(), "/tds_kv_", test_name); + std::filesystem::remove_all(test_dir_); + std::filesystem::create_directories(test_dir_); + } + + void TearDown() override { std::filesystem::remove_all(test_dir_); } + + std::string test_dir_; +}; + +// Verifies that page-aligned multi-slice buffers are written and read back +// end-to-end through `TdsKVBackend` (`tds_writev` / `tds_readv` with raw +// buffer descriptors, using `O_DIRECT` when supported by the test filesystem). +TEST_F(TdsKVBackendTest, DirectZeroCopyMultiSliceReadWrite) { + TdsKVBackend backend("tds_direct", {{"root_dir", test_dir_}, + {"storage_io_thread_pool_size", "4"}, + {"direct_io", "true"}}); + + const size_t align = GetStorageDirectIOAlignment(); + const size_t kSlice0Size = align; + const size_t kSlice1Size = 2 * align; + const size_t kSlice2Size = align; + const size_t kTotalBytes = kSlice0Size + kSlice1Size + kSlice2Size; + + AlignedBuffer src0(kSlice0Size); + AlignedBuffer src1(kSlice1Size); + AlignedBuffer src2(kSlice2Size); + ASSERT_NE(src0.data(), nullptr); + ASSERT_NE(src1.data(), nullptr); + ASSERT_NE(src2.data(), nullptr); + + std::vector expected0 = MakeDeterministicPayload(kSlice0Size, 1); + std::vector expected1 = MakeDeterministicPayload(kSlice1Size, 2); + std::vector expected2 = MakeDeterministicPayload(kSlice2Size, 3); + std::memcpy(src0.data(), expected0.data(), kSlice0Size); + std::memcpy(src1.data(), expected1.data(), kSlice1Size); + std::memcpy(src2.data(), expected2.data(), kSlice2Size); + + std::vector write_slices = { + HostBufferDescriptor{.ptr = src0.data(), .size = kSlice0Size}, + HostBufferDescriptor{.ptr = src1.data(), .size = kSlice1Size}, + HostBufferDescriptor{.ptr = src2.data(), .size = kSlice2Size}, + }; + + BlockKey key; + key.block_hash = "direct_block_01"; + key.resolved_key = absl::StrCat(test_dir_, "/direct_block_01.bin"); + key.offset = 0; + key.size = static_cast(kTotalBytes); + + absl::Notification write_done; + absl::Status write_status; + backend.WriteAsync(key, write_slices, kTotalBytes, [&](absl::Status status) { + write_status = std::move(status); + write_done.Notify(); + }); + write_done.WaitForNotification(); + ABSL_ASSERT_OK(write_status); + + AlignedBuffer dst0(kSlice0Size); + AlignedBuffer dst1(kSlice1Size); + AlignedBuffer dst2(kSlice2Size); + std::memset(dst0.data(), 0, kSlice0Size); + std::memset(dst1.data(), 0, kSlice1Size); + std::memset(dst2.data(), 0, kSlice2Size); + + std::vector read_slices = { + HostBufferDescriptor{.ptr = dst0.data(), .size = kSlice0Size}, + HostBufferDescriptor{.ptr = dst1.data(), .size = kSlice1Size}, + HostBufferDescriptor{.ptr = dst2.data(), .size = kSlice2Size}, + }; + + absl::Notification read_done; + absl::Status read_status; + backend.ReadAsync(key, read_slices, kTotalBytes, [&](absl::Status status) { + read_status = std::move(status); + read_done.Notify(); + }); + read_done.WaitForNotification(); + ABSL_ASSERT_OK(read_status); + + EXPECT_EQ(std::memcmp(dst0.data(), expected0.data(), kSlice0Size), 0); + EXPECT_EQ(std::memcmp(dst1.data(), expected1.data(), kSlice1Size), 0); + EXPECT_EQ(std::memcmp(dst2.data(), expected2.data(), kSlice2Size), 0); +} + +// Verifies that unaligned slice pointers/sizes automatically omit `O_DIRECT` +// when opening the file and complete multi-slice writes and reads via buffered +// `libtdsul` I/O without corrupting file size or payload bytes. +TEST_F(TdsKVBackendTest, UnalignedBufferedFallbackMultiSliceReadWrite) { + TdsKVBackend backend("tds_unaligned", {{"root_dir", test_dir_}, + {"storage_io_thread_pool_size", "4"}, + {"direct_io", "true"}}); + + constexpr size_t kSlice0Size = 1000; + constexpr size_t kSlice1Size = 513; + constexpr size_t kSlice2Size = 2000; + constexpr size_t kTotalBytes = kSlice0Size + kSlice1Size + kSlice2Size; + static_assert(kTotalBytes == 3513); + + std::vector src0 = MakeDeterministicPayload(kSlice0Size, 11); + std::vector src1 = MakeDeterministicPayload(kSlice1Size, 22); + std::vector src2 = MakeDeterministicPayload(kSlice2Size, 33); + + std::vector write_slices = { + HostBufferDescriptor{.ptr = src0.data(), .size = kSlice0Size}, + HostBufferDescriptor{.ptr = src1.data(), .size = kSlice1Size}, + HostBufferDescriptor{.ptr = src2.data(), .size = kSlice2Size}, + }; + + BlockKey key; + key.block_hash = "bounced_block_01"; + key.resolved_key = absl::StrCat(test_dir_, "/bounced_block_01.bin"); + key.offset = 0; + key.size = kTotalBytes; + + absl::Notification write_done; + absl::Status write_status; + backend.WriteAsync(key, write_slices, kTotalBytes, [&](absl::Status status) { + write_status = std::move(status); + write_done.Notify(); + }); + write_done.WaitForNotification(); + ABSL_ASSERT_OK(write_status); + + ASSERT_TRUE(std::filesystem::exists(key.resolved_key)); + EXPECT_EQ(std::filesystem::file_size(key.resolved_key), kTotalBytes); + + std::vector dst0(kSlice0Size, 0); + std::vector dst1(kSlice1Size, 0); + std::vector dst2(kSlice2Size, 0); + + std::vector read_slices = { + HostBufferDescriptor{.ptr = dst0.data(), .size = kSlice0Size}, + HostBufferDescriptor{.ptr = dst1.data(), .size = kSlice1Size}, + HostBufferDescriptor{.ptr = dst2.data(), .size = kSlice2Size}, + }; + + absl::Notification read_done; + absl::Status read_status; + backend.ReadAsync(key, read_slices, kTotalBytes, [&](absl::Status status) { + read_status = std::move(status); + read_done.Notify(); + }); + read_done.WaitForNotification(); + ABSL_ASSERT_OK(read_status); + + EXPECT_EQ(dst0, src0); + EXPECT_EQ(dst1, src1); + EXPECT_EQ(dst2, src2); +} + +// Verifies that `BatchExistsAsync` concurrently checks file existence across +// present and missing block keys and preserves input ordering in results. +TEST_F(TdsKVBackendTest, BatchExistsAsync) { + TdsKVBackend backend( + "tds_exists", + {{"root_dir", test_dir_}, {"storage_io_thread_pool_size", "8"}}); + + constexpr int kNumKeys = 16; + std::vector payload = MakeDeterministicPayload(2048, 7); + HostBufferDescriptor slice{.ptr = payload.data(), .size = payload.size()}; + + std::vector keys; + for (int i = 0; i < kNumKeys; ++i) { + BlockKey key; + key.block_hash = absl::StrCat("block_", i); + key.resolved_key = absl::StrCat(test_dir_, "/exists_", i, ".bin"); + key.offset = 0; + key.size = static_cast(payload.size()); + keys.push_back(key); + + if (i % 2 == 0) { + absl::Notification write_done; + absl::Status write_status; + backend.WriteAsync(key, {slice}, payload.size(), [&](absl::Status s) { + write_status = std::move(s); + write_done.Notify(); + }); + write_done.WaitForNotification(); + ABSL_ASSERT_OK(write_status); + } + } + + absl::Notification batch_done; + std::vector> results; + backend.BatchExistsAsync(keys, [&](std::vector> res) { + results = std::move(res); + batch_done.Notify(); + }); + batch_done.WaitForNotification(); + + ASSERT_EQ(results.size(), kNumKeys); + for (int i = 0; i < kNumKeys; ++i) { + EXPECT_THAT(results[i], IsOkAndHolds(i % 2 == 0)) << "Failed at key " << i; + } +} + +// Verifies that non-contiguous page-aligned sub-slices carved out of a larger +// contiguous host buffer pool are gathered on write and read back into a single +// contiguous destination buffer in slice order. +TEST_F(TdsKVBackendTest, SubSlicesOfAlignedPoolReadWrite) { + TdsKVBackend backend("tds_pool", {{"root_dir", test_dir_}}); + + const size_t align = GetStorageDirectIOAlignment(); + const size_t kPoolSize = 16 * align; + AlignedBuffer pool(kPoolSize); + ASSERT_NE(pool.data(), nullptr); + std::vector fill = MakeDeterministicPayload(kPoolSize, 55); + std::memcpy(pool.data(), fill.data(), kPoolSize); + + std::vector slices = { + HostBufferDescriptor{.ptr = pool.data() + align, .size = 2 * align}, + HostBufferDescriptor{.ptr = pool.data() + 8 * align, .size = align}, + }; + const size_t kTotalBytes = 3 * align; + BlockKey key; + key.block_hash = "pool_block"; + key.resolved_key = absl::StrCat(test_dir_, "/pool_block.bin"); + key.offset = 0; + key.size = static_cast(kTotalBytes); + + absl::Notification write_done; + absl::Status write_status; + backend.WriteAsync(key, slices, kTotalBytes, [&](absl::Status s) { + write_status = std::move(s); + write_done.Notify(); + }); + write_done.WaitForNotification(); + ABSL_ASSERT_OK(write_status); + + AlignedBuffer read_buf(kTotalBytes); + std::memset(read_buf.data(), 0, kTotalBytes); + absl::Notification read_done; + absl::Status read_status; + backend.ReadAsync( + key, {HostBufferDescriptor{.ptr = read_buf.data(), .size = kTotalBytes}}, + kTotalBytes, [&](absl::Status s) { + read_status = std::move(s); + read_done.Notify(); + }); + read_done.WaitForNotification(); + ABSL_ASSERT_OK(read_status); + EXPECT_EQ(std::memcmp(read_buf.data(), pool.data() + align, 2 * align), 0); + EXPECT_EQ(std::memcmp(read_buf.data() + 2 * align, pool.data() + 8 * align, + align), + 0); +} + +// Verifies that `TdsKVBackend` uses `StoragePathMapper` for shard-aware path +// mapping and is registered under `"tds"` in `KVCacheStoreBackendFactory`. +TEST_F(TdsKVBackendTest, + StoragePathMapperAndKVCacheStoreBackendFactoryIntegration) { + absl::flat_hash_map props = { + {"root_dir", test_dir_}, + {"model_name", "gemini_ultra"}, + {"tp_size", "8"}, + {"tp_rank", "3"}, + }; + TdsKVBackend backend("tds_mapper", props); + ASSERT_NE(backend.mapper(), nullptr); + EXPECT_EQ(backend.mapper()->tp_size(), 8); + + const std::string binary_hash("\xab\xcd\x00\x2f\xef\x01", 6); + absl::StatusOr mapped = backend.mapper()->MapKey(binary_hash); + ABSL_ASSERT_OK(mapped); + EXPECT_EQ(mapped->block_hash, binary_hash); + EXPECT_EQ( + mapped->resolved_key, + absl::StrCat(test_dir_, "/gemini_ultra/tp8_r3/abc/d0/abcd002fef01.bin")); + + BackendConfig store_cfg; + store_cfg.type = std::string(kTdsBackendName); + store_cfg.properties = { + {"root_dir", test_dir_}, + {"model_name", "gemini_ultra"}, + {"tp_size", "1"}, + {"capacity_bytes", "1048576"}, + }; + absl::StatusOr> store_backend = + KVCacheStoreBackendFactory::Instance().Create(store_cfg, + /*controller=*/nullptr); + ABSL_ASSERT_OK(store_backend); + ASSERT_NE(*store_backend, nullptr); + EXPECT_EQ((*store_backend)->name(), kTdsBackendName); + EXPECT_EQ((*store_backend)->GetCapacity(), 1048576); + + TdsKVBackend rank0_writer("tds_rank0", {{"root_dir", test_dir_}, + {"model_name", "gemini_ultra"}, + {"tp_size", "1"}, + {"tp_rank", "0"}}); + absl::StatusOr rank0_key = + rank0_writer.mapper()->MapKey(binary_hash); + ABSL_ASSERT_OK(rank0_key); + + const size_t align = GetStorageDirectIOAlignment(); + AlignedBuffer buf(align); + std::memset(buf.data(), 0x5A, align); + absl::Notification write_done; + absl::Status write_status; + rank0_writer.WriteAsync( + *rank0_key, {HostBufferDescriptor{.ptr = buf.data(), .size = align}}, + align, [&](absl::Status s) { + write_status = std::move(s); + write_done.Notify(); + }); + write_done.WaitForNotification(); + ABSL_ASSERT_OK(write_status); + + absl::StatusOr lookup_hits = + (*store_backend)->Lookup({binary_hash, "missing_hash"}); + ABSL_ASSERT_OK(lookup_hits); + ASSERT_EQ(lookup_hits->size(), 1); + EXPECT_EQ((*lookup_hits)[0].first, binary_hash); + EXPECT_EQ((*lookup_hits)[0].second.status, BlockStatus::SHARED_STORAGE); +} + +// Verifies concurrent multi-block, multi-slice writes followed by controller +// store lookup and concurrent reads across worker threads. +TEST_F(TdsKVBackendTest, ConcurrentMultiSliceWriteLookupAndRead) { + constexpr int kNumBlocks = 12; + const size_t kSliceBytes = GetStorageDirectIOAlignment(); + constexpr size_t kSlicesPerBlock = 3; + const size_t kBlockBytes = kSliceBytes * kSlicesPerBlock; + + BackendConfig store_cfg; + store_cfg.type = std::string(kTdsBackendName); + store_cfg.properties = { + {"root_dir", test_dir_}, + {"model_name", "concurrent_model"}, + {"tp_size", "1"}, + {"storage_io_thread_pool_size", "8"}, + }; + absl::StatusOr> store_backend = + KVCacheStoreBackendFactory::Instance().Create(store_cfg, + /*controller=*/nullptr); + ABSL_ASSERT_OK(store_backend); + + TdsKVBackend worker_backend("tds_worker", store_cfg.properties); + + std::vector> write_arena; + std::vector block_hashes; + std::vector mapped_keys; + for (int b = 0; b < kNumBlocks; ++b) { + block_hashes.push_back(absl::StrCat("concurrent_block_hash_", b)); + absl::StatusOr key = + worker_backend.mapper()->MapKey(block_hashes.back()); + ABSL_ASSERT_OK(key); + mapped_keys.push_back(*std::move(key)); + for (size_t s = 0; s < kSlicesPerBlock; ++s) { + auto buf = std::make_unique(kSliceBytes); + std::vector payload = MakeDeterministicPayload( + kSliceBytes, static_cast(b * 7 + s + 1)); + std::memcpy(buf->data(), payload.data(), kSliceBytes); + write_arena.push_back(std::move(buf)); + } + } + + std::atomic remaining_writes{kNumBlocks}; + absl::Notification all_writes_done; + std::vector write_statuses(kNumBlocks); + for (int b = 0; b < kNumBlocks; ++b) { + std::vector slices; + for (size_t s = 0; s < kSlicesPerBlock; ++s) { + slices.push_back(HostBufferDescriptor{ + .ptr = write_arena[b * kSlicesPerBlock + s]->data(), + .size = kSliceBytes}); + } + worker_backend.WriteAsync( + mapped_keys[b], slices, kBlockBytes, [&, b](absl::Status status) { + write_statuses[b] = std::move(status); + if (remaining_writes.fetch_sub(1, std::memory_order_acq_rel) == 1) { + all_writes_done.Notify(); + } + }); + } + all_writes_done.WaitForNotification(); + for (int b = 0; b < kNumBlocks; ++b) { + ABSL_ASSERT_OK(write_statuses[b]); + } + + absl::StatusOr hits = (*store_backend)->Lookup(block_hashes); + ABSL_ASSERT_OK(hits); + ASSERT_EQ(hits->size(), kNumBlocks); + for (int b = 0; b < kNumBlocks; ++b) { + EXPECT_EQ((*hits)[b].first, block_hashes[b]); + EXPECT_EQ((*hits)[b].second.status, BlockStatus::SHARED_STORAGE); + } + + std::vector> read_arena; + for (size_t i = 0; i < kNumBlocks * kSlicesPerBlock; ++i) { + auto buf = std::make_unique(kSliceBytes); + std::memset(buf->data(), 0, kSliceBytes); + read_arena.push_back(std::move(buf)); + } + + std::atomic remaining_reads{kNumBlocks}; + absl::Notification all_reads_done; + std::vector read_statuses(kNumBlocks); + for (int b = 0; b < kNumBlocks; ++b) { + std::vector slices; + for (size_t s = 0; s < kSlicesPerBlock; ++s) { + slices.push_back(HostBufferDescriptor{ + .ptr = read_arena[b * kSlicesPerBlock + s]->data(), + .size = kSliceBytes}); + } + worker_backend.ReadAsync( + mapped_keys[b], slices, kBlockBytes, [&, b](absl::Status status) { + read_statuses[b] = std::move(status); + if (remaining_reads.fetch_sub(1, std::memory_order_acq_rel) == 1) { + all_reads_done.Notify(); + } + }); + } + all_reads_done.WaitForNotification(); + for (int b = 0; b < kNumBlocks; ++b) { + ABSL_ASSERT_OK(read_statuses[b]); + for (size_t s = 0; s < kSlicesPerBlock; ++s) { + const size_t idx = b * kSlicesPerBlock + s; + EXPECT_EQ(std::memcmp(read_arena[idx]->data(), write_arena[idx]->data(), + kSliceBytes), + 0); + } + } +} + +// Verifies error status reporting for missing files, slice size mismatches, +// non-zero write offsets, and null slice pointers. +TEST_F(TdsKVBackendTest, ErrorHandling) { + TdsKVBackend backend("tds_errors", {{"root_dir", test_dir_}}); + + const size_t align = GetStorageDirectIOAlignment(); + AlignedBuffer buf(align); + + // 1. Reading a non-existent file returns kNotFound. + { + absl::Notification done; + absl::Status status; + backend.ReadAsync( + BlockKey{.block_hash = "non_existent", + .resolved_key = absl::StrCat(test_dir_, "/missing.bin"), + .offset = 0, + .size = static_cast(align)}, + {HostBufferDescriptor{.ptr = buf.data(), .size = align}}, align, + [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + EXPECT_THAT(status, StatusIs(absl::StatusCode::kNotFound)); + } + + // 2. Writing with mismatched total_bytes returns kInvalidArgument. + { + absl::Notification done; + absl::Status status; + backend.WriteAsync( + BlockKey{.block_hash = "mismatch", + .resolved_key = absl::StrCat(test_dir_, "/mismatch.bin"), + .offset = 0, + .size = static_cast(2 * align)}, + {HostBufferDescriptor{.ptr = buf.data(), .size = align}}, 2 * align, + [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + EXPECT_THAT(status, StatusIs(absl::StatusCode::kInvalidArgument)); + } + + // 3. Writing with non-zero offset returns kInvalidArgument. + { + absl::Notification done; + absl::Status status; + backend.WriteAsync( + BlockKey{.block_hash = "nonzero_offset", + .resolved_key = absl::StrCat(test_dir_, "/nonzero.bin"), + .offset = static_cast(align), + .size = static_cast(align)}, + {HostBufferDescriptor{.ptr = buf.data(), .size = align}}, align, + [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + EXPECT_THAT(status, StatusIs(absl::StatusCode::kInvalidArgument)); + } + + // 4. Null slice pointer with non-zero size returns kInvalidArgument. + { + absl::Notification done; + absl::Status status; + backend.WriteAsync( + BlockKey{.block_hash = "null_slice", + .resolved_key = absl::StrCat(test_dir_, "/null_slice.bin"), + .offset = 0, + .size = static_cast(align)}, + {HostBufferDescriptor{.ptr = nullptr, .size = align}}, align, + [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + EXPECT_THAT(status, StatusIs(absl::StatusCode::kInvalidArgument)); + } +} + +// Verifies that reading from an unaligned file offset falls back to buffered +// I/O (omitting `O_DIRECT`) and returns the exact sub-range of bytes at +// `key.offset`. +TEST_F(TdsKVBackendTest, UnalignedOffsetReadReturnsCorrectBytes) { + TdsKVBackend backend("tds_unaligned_read", + {{"root_dir", test_dir_}, + {"storage_io_thread_pool_size", "2"}, + {"direct_io", "true"}}); + + const size_t align = GetStorageDirectIOAlignment(); + const size_t kBaselineBytes = 3 * align; + AlignedBuffer src(kBaselineBytes); + ASSERT_NE(src.data(), nullptr); + const std::vector baseline = + MakeDeterministicPayload(kBaselineBytes, 97); + std::memcpy(src.data(), baseline.data(), kBaselineBytes); + + BlockKey write_key; + write_key.block_hash = "bounce_read_baseline"; + write_key.resolved_key = absl::StrCat(test_dir_, "/bounce_read.bin"); + write_key.offset = 0; + write_key.size = static_cast(kBaselineBytes); + + { + absl::Notification done; + absl::Status status; + backend.WriteAsync( + write_key, + {HostBufferDescriptor{.ptr = src.data(), .size = kBaselineBytes}}, + kBaselineBytes, [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + ABSL_ASSERT_OK(status); + } + + const int64_t kReadOffset = static_cast(align) + 904; + constexpr size_t kReadBytes = 300; + std::vector dst(kReadBytes, 0); + + BlockKey read_key; + read_key.block_hash = "bounce_read_baseline"; + read_key.resolved_key = write_key.resolved_key; + read_key.offset = kReadOffset; + read_key.size = kReadBytes; + + { + absl::Notification done; + absl::Status status; + backend.ReadAsync( + read_key, {HostBufferDescriptor{.ptr = dst.data(), .size = kReadBytes}}, + kReadBytes, [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + ABSL_ASSERT_OK(status); + } + + EXPECT_EQ(std::memcmp(dst.data(), baseline.data() + kReadOffset, kReadBytes), + 0); +} + +// Verifies that scatter-gather transfers with more than `UIO_MAXIOV` (1024) +// slices are split into multiple `tds_writev` / `tds_readv` batches with +// advancing file offsets. +TEST_F(TdsKVBackendTest, ScatterGatherCrossesIovLimitInMultipleBatches) { + TdsKVBackend backend("tds_iov", + {{"root_dir", test_dir_}, + {"storage_io_thread_pool_size", "2"}, + {"direct_io", "true"}}); + + const size_t align = GetStorageDirectIOAlignment(); + constexpr size_t kPages = static_cast(UIO_MAXIOV) + 1; + const size_t kTotalBytes = kPages * align; + + AlignedBuffer src(kTotalBytes); + AlignedBuffer dst(kTotalBytes); + ASSERT_NE(src.data(), nullptr); + ASSERT_NE(dst.data(), nullptr); + + const std::vector payload = + MakeDeterministicPayload(kTotalBytes, 53); + std::memcpy(src.data(), payload.data(), kTotalBytes); + std::memset(dst.data(), 0, kTotalBytes); + + std::vector write_slices; + std::vector read_slices; + write_slices.reserve(kPages); + read_slices.reserve(kPages); + for (size_t i = 0; i < kPages; ++i) { + write_slices.push_back( + HostBufferDescriptor{.ptr = src.data() + i * align, .size = align}); + read_slices.push_back( + HostBufferDescriptor{.ptr = dst.data() + i * align, .size = align}); + } + + BlockKey key; + key.block_hash = "iov_block"; + key.resolved_key = absl::StrCat(test_dir_, "/iov_block.bin"); + key.offset = 0; + key.size = static_cast(kTotalBytes); + + { + absl::Notification done; + absl::Status status; + backend.WriteAsync(key, write_slices, kTotalBytes, [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + ABSL_ASSERT_OK(status); + } + + EXPECT_EQ(FileSizeOrNegative(key.resolved_key), + static_cast(kTotalBytes)); + + { + absl::Notification done; + absl::Status status; + backend.ReadAsync(key, read_slices, kTotalBytes, [&](absl::Status s) { + status = std::move(s); + done.Notify(); + }); + done.WaitForNotification(); + ABSL_ASSERT_OK(status); + } + + EXPECT_EQ(std::memcmp(dst.data(), src.data(), kTotalBytes), 0); +} + +} // namespace +} // namespace storage +} // namespace backends +} // namespace kv_cache +} // namespace tpu_raiden diff --git a/tpu_sync/kv_cache/kv_cache_manager_base.cc b/tpu_sync/kv_cache/kv_cache_manager_base.cc index 40055e6c4..2b6158c1b 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_base.cc +++ b/tpu_sync/kv_cache/kv_cache_manager_base.cc @@ -66,6 +66,7 @@ #include "tpu_sync/core/tpu_utils.h" #include "tpu_sync/kv_cache/backends/backend.h" #include "tpu_sync/kv_cache/backends/storage/posix_backend.h" +#include "tpu_sync/kv_cache/backends/storage/tds_backend.h" #include "tpu_sync/kv_cache/kv_cache_store_backend_factory.h" #include "tpu_sync/kv_cache/logical_block_manager.h" #include "tpu_sync/kv_cache/pool_layout.h" @@ -3600,8 +3601,11 @@ bool KVCacheManagerBase::InitializeSingleSecondaryBackend( const BackendConfig& config) { if (config.type.empty()) return false; - if (!absl::EqualsIgnoreCase(config.type, - backends::storage::kPosixBackendName)) { + const bool is_posix = + absl::EqualsIgnoreCase(config.type, backends::storage::kPosixBackendName); + const bool is_tds = + absl::EqualsIgnoreCase(config.type, backends::storage::kTdsBackendName); + if (!is_posix && !is_tds) { LOG(WARNING) << "[Worker] Unsupported secondary backend: " << config.type; return false; } @@ -3613,7 +3617,8 @@ bool KVCacheManagerBase::InitializeSingleSecondaryBackend( } const std::string canonical_name = - std::string(backends::storage::kPosixBackendName); + std::string(is_tds ? backends::storage::kTdsBackendName + : backends::storage::kPosixBackendName); if (GetKVBackend(canonical_name) != nullptr) return false; // Storage topology comes only from BackendConfig::parallelism, resolved the @@ -3624,17 +3629,26 @@ bool KVCacheManagerBase::InitializeSingleSecondaryBackend( config.parallelism.tp_size > 0 ? config.parallelism.tp_size : 1, .tp_rank = config.parallelism.tp_rank}; ApplyParallelismToProperties(effective, &resolved); - if (absl::Status status = - backends::storage::PosixBackendOptions::FromProperties( - resolved.properties) - .status(); - !status.ok()) { + const absl::Status status = + is_tds ? backends::storage::TdsBackendOptions::FromProperties( + resolved.properties) + .status() + : backends::storage::PosixBackendOptions::FromProperties( + resolved.properties) + .status(); + if (!status.ok()) { LOG(ERROR) << "[Worker] invalid " << config.type << " backend config; refusing to register: " << status; return false; } - auto backend = std::make_shared( - canonical_name, resolved.properties); + std::shared_ptr backend; + if (is_tds) { + backend = std::make_shared( + canonical_name, resolved.properties); + } else { + backend = std::make_shared( + canonical_name, resolved.properties); + } { absl::MutexLock lock(backends_mu_); backends_[canonical_name] = std::move(backend); diff --git a/tpu_sync/kv_cache/kv_cache_store_test.cc b/tpu_sync/kv_cache/kv_cache_store_test.cc index 3565f58c7..6dfc4f4a0 100644 --- a/tpu_sync/kv_cache/kv_cache_store_test.cc +++ b/tpu_sync/kv_cache/kv_cache_store_test.cc @@ -71,6 +71,7 @@ #include "tpu_sync/core/kv_manager_holder.h" #include "tpu_sync/core/raiden_transfer_endpoint.h" #include "tpu_sync/kv_cache/backends/backend.h" +#include "tpu_sync/kv_cache/backends/storage/tds_backend.h" #include "tpu_sync/kv_cache/block_tracker.h" #include "tpu_sync/kv_cache/global_registry/global_registry.grpc.pb.h" #include "tpu_sync/kv_cache/global_registry/global_registry_client.h" @@ -6537,12 +6538,17 @@ TEST(KVCacheStoreTest, MultiBackendConstructionValidation) { // Accepts 1 tier-0 backend (0 secondary). EXPECT_TRUE(CreateStore(std::vector{tier0}).ok()); - // Accepts 1 tier-0 backend + 1 secondary backend. + // Accepts 1 tier-0 backend + 1 secondary backend ("posix" or "tds"). BackendConfig sec1; sec1.type = "posix"; sec1.SetProperty("storage_root", "/tmp/raiden_test_sec1"); EXPECT_TRUE(CreateStore(std::vector{tier0, sec1}).ok()); + BackendConfig tds_sec; + tds_sec.type = "tds"; + tds_sec.SetProperty("root_dir", "/tmp/raiden_test_tds"); + EXPECT_TRUE(CreateStore(std::vector{tier0, tds_sec}).ok()); + // Rejects 1 tier-0 backend + 2 secondary backends. BackendConfig sec2; sec2.type = "posix"; @@ -6620,6 +6626,82 @@ TEST_F(KVCacheStoreEmbeddedControllerTest, EXPECT_GE((*backend_lookup)[0].second.host_block_id, 0); } +// ----------------------------------------------------------------------------- +// Verifies the RECALL path (`KVCacheStore::Load`) for a block held by the +// `"tds"` secondary backend (storage -> [dram: User host_buf] -> [hbm]): +// 1. `KVCacheStore` allocates a staging block in [dram: User host_buf] and +// dispatches the recall to the worker's `MockTransferManager`, which +// records one H2D (`h2d_calls`); no storage I/O is performed. +// 2. After `PollLoadStatus()` reports the hash in `load_done`, the store's +// lookup reports `BlockStatus::HOST_AND_HBM` with `device_block_id == 4` +// and a valid `host_block_id`, i.e. the staging block is kept. +// ----------------------------------------------------------------------------- +TEST_F(KVCacheStoreEmbeddedControllerTest, + TdsStorageRecallRetainsStagingHostBlocksAsHostAndHbm) { + ::tpu_raiden::controller::MockTransferManager mock_mgr; + test_server_->service->SetTransferManager( + ::tpu_raiden::KVManagerHolder(&mock_mgr)); + + auto controller = MakeController(); + RegisterAndInitWorker(*controller, "worker_0", test_server_->server_address); + + RaidenId rid{"test_job", "0", "test_cache", 0}; + KVCacheStore store(10, std::move(controller), "", rid, std::nullopt, + /*store_server_ip=*/"127.0.0.1"); + + // Step 1: Attach a `"tds"` (`TdsKVCacheStoreBackend` + `TdsKVBackend`) + // secondary backend to both the Coordinator store and the Worker manager. + std::string scratch_dir = + std::string(testing::TempDir()) + "/" + + ::testing::UnitTest::GetInstance()->current_test_info()->name(); + auto worker_backend = std::make_shared( + "tds", absl::flat_hash_map{ + {"root_dir", scratch_dir}, + {"model_name", "model_test"}, + {"tp_size", "1"}, + {"tp_rank", "0"}}); + KVCacheStoreTest::AddBackend( + store, std::make_shared( + worker_backend, "tds")); + mock_mgr.backends["tds"] = worker_backend; + + // Step 2: Request `store.Load()` for a `BlockStatus::SHARED_STORAGE` slice + // residing in `"tds"` into `[hbm]` block 4. + RaidenBlockId slice(RaidenId{"test_job", "0", "tds", 0}, -1, -1, + BlockStatus::SHARED_STORAGE); + + absl::Status status = store.Load({"tds_storage_hash"}, {slice}, {4}); + ABSL_ASSERT_OK(status); + + // Step 3: Poll until the async recall finishes and verify the block is now + // tracked as `BlockStatus::HOST_AND_HBM` (`device_block_id == 4` and a valid + // `host_block_id >= 0` in `[dram: User host_buf]`). + bool done = false; + while (!done) { + auto [load_done, load_failed, load_pending, load_existing, + load_unregistered] = store.PollLoadStatus(); + if (!load_failed.empty()) { + FAIL() << "Async Load failed during polling: " << load_failed[0]; + } + if (!load_done.empty()) { + EXPECT_THAT(load_done, ::testing::ElementsAre("tds_storage_hash")); + done = true; + } + if (!done) { + absl::SleepFor(absl::Milliseconds(10)); + } + } + + EXPECT_EQ(mock_mgr.h2d_calls, 1); + + auto lookup_res = PeekLookup(store, {"tds_storage_hash"}); + ABSL_ASSERT_OK(lookup_res); + ASSERT_EQ(lookup_res->size(), 1); + EXPECT_EQ((*lookup_res)[0].second.status, BlockStatus::HOST_AND_HBM); + EXPECT_EQ((*lookup_res)[0].second.device_block_id, 4); + EXPECT_GE((*lookup_res)[0].second.host_block_id, 0); +} + TEST_F(KVCacheStoreEmbeddedControllerTest, StorageRecallFailureDeallocatesStagingHostBlocks) { ::tpu_raiden::controller::MockTransferManager mock_mgr; @@ -7049,29 +7131,110 @@ TEST_F(KVCacheStoreEmbeddedControllerTest, EXPECT_EQ((*lookup_res)[0].second.status, BlockStatus::HOST_AND_HBM); } +// ----------------------------------------------------------------------------- +// Verifies the OFFLOAD path (`KVCacheStore::Save`) for a block resident only in +// [hbm] (`BlockStatus::HBM`) with the `"tds"` secondary backend attached +// ([hbm] -> [dram: User host_buf] -> storage): +// 1. `KVCacheStore::Save` allocates a staging block in [dram: User host_buf] +// and dispatches the transfer to the worker's `MockTransferManager`, +// which records one D2H (`d2h_calls`); no storage I/O is performed. +// 2. After `PollSaveStatus()` reports the hash in `save_done` (with nothing +// in `save_existing` / `save_unregistered`), `d2h_calls == 1` and the +// store's lookup reports `BlockStatus::HOST_AND_HBM`. +// ----------------------------------------------------------------------------- +TEST_F(KVCacheStoreEmbeddedControllerTest, + TdsSaveLocalDispatchesOffloadToSecondaryBackend) { + ::tpu_raiden::controller::MockTransferManager mock_mgr; + test_server_->service->SetTransferManager( + ::tpu_raiden::KVManagerHolder(&mock_mgr)); + + auto controller = MakeController(); + RegisterAndInitWorker(*controller, "worker_0", test_server_->server_address); + + RaidenId rid{"test_job", "0", "test_cache", 0}; + KVCacheStore store(10, std::move(controller), "", rid, std::nullopt, + /*store_server_ip=*/"127.0.0.1"); + + // Step 1: Attach the `"tds"` secondary backend to the store and worker. + std::string scratch_dir = + std::string(testing::TempDir()) + "/" + + ::testing::UnitTest::GetInstance()->current_test_info()->name(); + auto worker_backend = std::make_shared( + "tds", absl::flat_hash_map{ + {"root_dir", scratch_dir}, + {"model_name", "model_test"}, + {"tp_size", "1"}, + {"tp_rank", "0"}}); + KVCacheStoreTest::AddBackend( + store, std::make_shared( + worker_backend, "tds")); + mock_mgr.backends["tds"] = worker_backend; + + // Step 2: Seed `"tds_save_hash"` as resident in `[hbm]` (`device_block_id == + // 0`, `BlockStatus::HBM`) and invoke `store.Save()`. + std::vector hashes = {"tds_save_hash"}; + std::vector slices = { + RaidenBlockId(rid, -1, 0, BlockStatus::HBM)}; + + ASSERT_TRUE(InsertResident(store, hashes, slices, /*on_host=*/false)); + ABSL_ASSERT_OK(store.Lookup(hashes)); + + absl::Status status = store.Save(hashes); + ABSL_ASSERT_OK(status); + + // Step 3: Poll until `PollSaveStatus()` completes and verify the block is + // now `BlockStatus::HOST_AND_HBM` with `d2h_calls == 1`. + bool done = false; + while (!done) { + auto [save_done, save_failed, save_pending, save_existing, + save_unregistered] = store.PollSaveStatus(); + if (!save_failed.empty()) { + FAIL() << "Async Save failed during polling: " << save_failed[0]; + } + EXPECT_TRUE(save_existing.empty()); + EXPECT_TRUE(save_unregistered.empty()); + if (!save_done.empty()) { + EXPECT_THAT(save_done, ::testing::ElementsAre("tds_save_hash")); + done = true; + } + if (!done) { + absl::SleepFor(absl::Milliseconds(10)); + } + } + + EXPECT_EQ(mock_mgr.d2h_calls, 1); + + auto lookup_res = PeekLookup(store, hashes); + ABSL_ASSERT_OK(lookup_res); + ASSERT_EQ(lookup_res->size(), 1); + EXPECT_EQ((*lookup_res)[0].second.status, BlockStatus::HOST_AND_HBM); +} + TEST(KVCacheStoreTest, CreateWithProgrammaticSecondaryConfigs) { BackendConfig host_cfg; host_cfg.type = "HostOffloadBackend"; host_cfg.capacity = 4; - BackendConfig sec_cfg; - sec_cfg.type = "posix"; - sec_cfg.properties["root_dir"] = "/tmp/test"; + for (absl::string_view backend_type : {"posix", "tds"}) { + BackendConfig sec_cfg; + sec_cfg.type = std::string(backend_type); + sec_cfg.properties["root_dir"] = "/tmp/test"; - std::vector sec_cfgs = {sec_cfg}; - auto store_or = KVCacheStore::Create( - host_cfg, /*capacity=*/4, /*global_registry_address=*/"", RaidenId{}, - /*num_shards=*/1, /*shard_size_bytes=*/512, - /*store_server_ip=*/"127.0.0.1", /*raiden_controller_port=*/0, - /*metadata=*/std::nullopt, /*expected_worker_count=*/0, - std::move(sec_cfgs)); - ASSERT_TRUE(store_or.ok()) << store_or.status(); - auto& store = *store_or; - ASSERT_EQ(store->backends().size(), 2); - EXPECT_EQ(store->backends()[0]->name(), "HostOffloadBackend"); - EXPECT_EQ(store->backends()[1]->name(), "posix"); - ASSERT_EQ(store->backend_configs().size(), 2); - EXPECT_EQ(store->backend_configs()[1].GetProperty("root_dir"), "/tmp/test"); + std::vector sec_cfgs = {sec_cfg}; + auto store_or = KVCacheStore::Create( + host_cfg, /*capacity=*/4, /*global_registry_address=*/"", RaidenId{}, + /*num_shards=*/1, /*shard_size_bytes=*/512, + /*store_server_ip=*/"127.0.0.1", /*raiden_controller_port=*/0, + /*metadata=*/std::nullopt, /*expected_worker_count=*/0, + std::move(sec_cfgs)); + ASSERT_TRUE(store_or.ok()) << store_or.status(); + auto& store = *store_or; + ASSERT_EQ(store->backends().size(), 2); + EXPECT_EQ(store->backends()[0]->name(), "HostOffloadBackend"); + EXPECT_EQ(store->backends()[1]->name(), backend_type); + ASSERT_EQ(store->backend_configs().size(), 2); + EXPECT_EQ(store->backend_configs()[1].GetProperty("root_dir"), "/tmp/test"); + } } TEST(KVCacheStoreTest, PeerLookupPriorityOverStorageFallback) { diff --git a/tpu_sync/tpudirect_storage/BUILD b/tpu_sync/tpudirect_storage/BUILD new file mode 100644 index 000000000..057ec6fc0 --- /dev/null +++ b/tpu_sync/tpudirect_storage/BUILD @@ -0,0 +1,190 @@ +# Copyright 2026 Google LLC. +# +# 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. + +load("@rules_cc//cc:defs.bzl", "cc_library", "cc_test") + +# TPUDirect Storage Userspace Library (`libtdsul`). +# +# Vendored copy of the upstream `tpudirect-storage` repository +# (https://cloudrdma-guest.git.corp.google.com/tpudirect-storage, snapshot of +# commit e046a94), which has no permanent google3 home yet. It lives inside +# `//third_party/tpu_raiden` so that `TdsKVBackend` can depend on it without the +# illegal non-experimental -> experimental edge that `//experimental/xpu-rdma/ +# tpudirect-storage:tdsul` produced (that edge builds only under +# `--experimental_deps_ok` and can never be submitted). +# +# TODO: replace this vendored copy with a dependency on the canonical +# `libtdsul` target once the TDSUL owners establish one in google3. +# +# Local modifications relative to upstream (re-apply these on any re-sync; a +# plain re-copy will silently drop them): +# 1. `src/util/ThreadPoolConductor.cpp` -- `ThreadPoolConductor::Start()` now +# pins workers to CPUs only when an explicit core range was configured. +# Upstream always pinned, deriving core indices from +# `sysconf(_SC_NPROCESSORS_ONLN)`, which reports machine-wide cores and +# says nothing about the process's allowed cpuset. Under a restricted +# cpuset every `pthread_create` failed with EINVAL. Any job restricted to +# a cpuset can hit this, not just test sandboxes. +# 2. `src/util/ThreadPoolConductor.cpp` -- `WorkerThread::Start()` retries +# unpinned when an affinity-pinned `pthread_create` fails, instead of +# logging and leaving a pool that accepts work it will never run. +# 3. `src/util/ThreadPoolConductor.cpp` -- `WorkerThread::Stop()` drains +# `task_queue_`, failing the futures with `TDS_ERROR_IO` and freeing +# requests it owns. Upstream leaked every request queued to a worker +# whose thread never started. +# `tests/thread_pool_conductor_test.cpp` is correspondingly strengthened to +# assert futures actually complete rather than only that nothing crashed. +# Upstream's `kokoro/` CI directory is intentionally not vendored. +# +# These three fixes have not yet been reported upstream. +# +# `libtdsul` is a zero-copy execution engine: it performs POSIX +# `pread`/`pwrite`/`preadv`/`pwritev` between [local-storage: NVMe SSD] / +# [shared-storage: Lustre/NFS] and `TDS_MEM_HOST` buffers (registered handles +# or raw pointers) in [dram: User host_buf]. Per `go/tdsul-api` §1 it holds no +# internal bounce buffer pools; alignment fallbacks belong to the framework +# adapter +# (`//tpu_sync/kv_cache/backends/storage:tds_backend`). +package(default_visibility = ["//visibility:public"]) + +cc_library( + name = "tdsul", + srcs = [ + "src/tdsul.cpp", + "src/util/IoQueue.cpp", + "src/util/IoUringConductor.cpp", + "src/util/ThreadPoolConductor.cpp", + ], + hdrs = [ + "include/tdsul/def.h", + "include/tdsul/tdsul.h", + "src/def_internal.h", + "src/syscall_internal.h", + "src/util/IoConductor.h", + "src/util/IoQueue.h", + "src/util/IoUringConductor.h", + "src/util/ThreadPoolConductor.h", + ], + # Upstream sources use bare `absl/...` includes rather than google3's + # `third_party/absl/...` convention, so `third_party` is added as a system + # include root instead of rewriting every upstream include. + copts = [ + "-isystem", + "third_party", + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + # Upstream sources include `absl/...`, `gtest/...` and `gmock/...` + # directly rather than through google3 module-exporting targets, which + # `layering_check` cannot resolve. Sibling tpu_raiden targets (e.g. + # `:posix_backend`) disable it for the same reason. + "-layering_check", + ], + includes = [ + "include", + "src", + ], + visibility = ["//visibility:public"], + deps = [ + "@com_google_absl//absl/log", + "@com_google_absl//absl/strings", + ], +) + +cc_test( + name = "tdsul_test", + srcs = ["tests/tdsul_test.cpp"], + copts = [ + "-isystem", + "third_party", + "-isystem", + "third_party/googletest/googletest/include", + "-isystem", + "third_party/googletest/googlemock/include", + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + # Upstream sources include `absl/...`, `gtest/...` and `gmock/...` + # directly rather than through google3 module-exporting targets, which + # `layering_check` cannot resolve. Sibling tpu_raiden targets (e.g. + # `:posix_backend`) disable it for the same reason. + "-layering_check", + ], + deps = [ + ":tdsul", + "@com_google_absl//absl/log", + "@com_google_absl//absl/log:initialize", + "@com_google_googletest//:gtest_main", + ], +) + +cc_test( + name = "thread_pool_conductor_test", + srcs = ["tests/thread_pool_conductor_test.cpp"], + copts = [ + "-isystem", + "third_party", + "-isystem", + "third_party/googletest/googletest/include", + "-isystem", + "third_party/googletest/googlemock/include", + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + # Upstream sources include `absl/...`, `gtest/...` and `gmock/...` + # directly rather than through google3 module-exporting targets, which + # `layering_check` cannot resolve. Sibling tpu_raiden targets (e.g. + # `:posix_backend`) disable it for the same reason. + "-layering_check", + ], + deps = [ + ":tdsul", + "@com_google_absl//absl/log", + "@com_google_absl//absl/log:initialize", + "@com_google_googletest//:gtest_main", + ], +) + +cc_test( + name = "io_queue_test", + srcs = ["tests/io_queue_test.cpp"], + copts = [ + "-isystem", + "third_party", + "-isystem", + "third_party/googletest/googletest/include", + "-isystem", + "third_party/googletest/googlemock/include", + "-fno-strict-aliasing", + "-fexceptions", + ], + features = [ + "-use_header_modules", + # Upstream sources include `absl/...`, `gtest/...` and `gmock/...` + # directly rather than through google3 module-exporting targets, which + # `layering_check` cannot resolve. Sibling tpu_raiden targets (e.g. + # `:posix_backend`) disable it for the same reason. + "-layering_check", + ], + deps = [ + ":tdsul", + "@com_google_googletest//:gtest_main", + ], +) diff --git a/tpu_sync/tpudirect_storage/include/tdsul/def.h b/tpu_sync/tpudirect_storage/include/tdsul/def.h new file mode 100644 index 000000000..d610a7001 --- /dev/null +++ b/tpu_sync/tpudirect_storage/include/tdsul/def.h @@ -0,0 +1,155 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_DEF_H_ +#define TDSUL_DEF_H_ + +#include +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +/* ========================================================================== */ +/* ENUMERATIONS */ +/* ========================================================================== */ + +typedef enum tds_result { + TDS_SUCCESS = 0, + TDS_PENDING = 1, + TDS_PARTIAL_FINISH = 2, + TDS_ERROR_INITIALIZATION = 3, + TDS_ERROR_NO_MEMORY = 4, + TDS_ERROR_INVALID_PARAMETER = 5, + TDS_ERROR_UNALIGNED_BUFFER = 6, + TDS_ERROR_P2P_UNSUPPORTED = 7, + TDS_ERROR_IN_USE = 8, + TDS_ERROR_IO = 9, + TDS_ERROR_UNSUPPORTED = 10 +} tds_result_t; + +typedef enum { TDS_MEM_HOST = 0, TDS_MEM_DEVICE = 1 } MemoryType; + +typedef enum { TDS_OP_READ = 0, TDS_OP_WRITE = 1 } OpType; + +/* ========================================================================== */ +/* OPAQUE HANDLES */ +/* ========================================================================== */ + +typedef struct tds_config tds_config_t; +typedef struct tds_storage_opts tds_storage_opts_t; +typedef struct tds_storage_handle tds_storage_handle_t; +typedef struct tds_buffer_handle tds_buffer_handle_t; +typedef struct tds_io_queue tds_io_queue_t; +typedef struct tds_io_future tds_future_t; +typedef struct tds_batch tds_batch_t; + +/* ========================================================================== */ +/* TRANSPARENT DESCRIPTORS */ +/* ========================================================================== */ + +typedef struct tds_storage_descr { + const char* uri; + tds_storage_opts_t* options; +} tds_storage_descr_t; + +struct tds_file_io_args { + off_t + file_offset; // Used by POSIX file and borrowed file descriptor backends + size_t size; // optional, using the buffer-side size as the ground truth. +}; + +typedef struct tds_storage_io { + tds_storage_handle_t* handle; + union { + struct tds_file_io_args file; + uint64_t reserved[7]; // 56 bytes reserved inside union (aligns struct to + // 64-byte cache line) + }; +} tds_storage_io_t; + +struct tds_raw_ptr_io_args { + void* vptr; + size_t size; +}; + +struct tds_registered_io_args { + size_t offset; + size_t size; +}; + +typedef struct tds_buffer_io { + tds_buffer_handle_t* + handle; // Registered handle, or NULL for on-the-fly resolution via vaddr + union { + struct tds_raw_ptr_io_args + raw_ptr; // Virtual address pointer (used when handle == NULL) + struct tds_registered_io_args + registered; // Slice offset in bytes (used when handle != NULL) + uint64_t reserved[7]; // 56 bytes reserved inside union (aligns struct to + // 64-byte cache line) + }; +} tds_buffer_io_t; + +typedef struct tds_io_status { + tds_result_t status; // e.g., TDS_SUCCESS, TDS_ERROR_IO, TDS_PARTIAL_FINISH + int error_num; // POSIX errno if an error occurred + size_t finished_bytes; // Bytes successfully transferred +} tds_io_status_t; + +/* TODO: Add GcsIoArgs and GCS storage backend support in the future */ + +/* ========================================================================== */ +/* CREATION HELPER INLINES */ +/* ========================================================================== */ + +static inline tds_storage_io_t tds_create_file_io(tds_storage_handle_t* handle, + off_t file_offset, + size_t size) { + tds_storage_io_t io; + memset(&io, 0, sizeof(io)); + io.handle = handle; + io.file.file_offset = file_offset; + io.file.size = size; + return io; +} + +static inline tds_buffer_io_t tds_create_registered_buffer_io( + tds_buffer_handle_t* handle, size_t offset, size_t size) { + tds_buffer_io_t io; + memset(&io, 0, sizeof(io)); + io.handle = handle; + io.registered.offset = offset; + io.registered.size = size; + return io; +} + +static inline tds_buffer_io_t tds_create_raw_buffer_io(void* vptr, + size_t size) { + tds_buffer_io_t io; + memset(&io, 0, sizeof(io)); + io.handle = NULL; + io.raw_ptr.vptr = vptr; + io.raw_ptr.size = size; + return io; +} + +#ifdef __cplusplus +} +#endif + +#endif // TDSUL_DEF_H_ diff --git a/tpu_sync/tpudirect_storage/include/tdsul/tdsul.h b/tpu_sync/tpudirect_storage/include/tdsul/tdsul.h new file mode 100644 index 000000000..2f271e1ba --- /dev/null +++ b/tpu_sync/tpudirect_storage/include/tdsul/tdsul.h @@ -0,0 +1,143 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_H_ +#define TDSUL_H_ + +#include +#include + +#include "tdsul/def.h" + +#ifdef __cplusplus +extern "C" { +#endif + +/* ========================================================================== */ +/* 4.1 Initialization & Lifecycle Management */ +/* ========================================================================== */ + +tds_config_t* tds_config_create(void); +void tds_config_destroy(tds_config_t* config); +tds_result_t tds_config_set_int(tds_config_t* config, const char* key, + int value); +tds_result_t tds_config_set_string(tds_config_t* config, const char* key, + const char* value); + +tds_storage_opts_t* tds_storage_opts_create(void); +void tds_storage_opts_destroy(tds_storage_opts_t* opts); +tds_result_t tds_storage_opts_set_int(tds_storage_opts_t* opts, const char* key, + int value); +tds_result_t tds_storage_opts_set_string(tds_storage_opts_t* opts, + const char* key, const char* value); + +tds_result_t tds_init(const tds_config_t* config); +tds_result_t tds_shutdown(void); + +/* ========================================================================== */ +/* 4.2 Registration Management */ +/* ========================================================================== */ + +tds_result_t tds_buffer_register_vaddr(void* vaddr, int memory_type, + size_t size, + tds_buffer_handle_t** buf_handle); +tds_result_t tds_buffer_register_dmabuf(int dma_buf_fd, size_t offset, + int memory_type, size_t size, + tds_buffer_handle_t** buf_handle); +tds_result_t tds_buffer_deregister(tds_buffer_handle_t* buf_handle); + +tds_result_t tds_storage_handle_register(const tds_storage_descr_t* descr, + tds_storage_handle_t** storage_handle); +tds_result_t tds_storage_handle_deregister( + tds_storage_handle_t* storage_handle); + +/* ========================================================================== */ +/* 4.3 Synchronous I/O */ +/* ========================================================================== */ + +ssize_t tds_read(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io); +ssize_t tds_write(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io); +/** + * @brief Synchronous vectored read/write operations. + * + * @note Unlike scalar I/O (tds_read/tds_write) which defaults a 0-length buffer + * to the size of the file mapping, vectored I/O treats a 0-length segment as a + * valid empty buffer and does NOT default it to the file size. This matches + * standard POSIX behavior to prevent unexpected buffer overflows when + * iterating. + */ +ssize_t tds_readv(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, int num_buffers); +ssize_t tds_writev(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, int num_buffers); + +/* ========================================================================== */ +/* 4.4 Asynchronous I/O & Queue Management */ +/* ========================================================================== */ + +tds_result_t tds_queue_create(tds_io_queue_t** queue); +tds_result_t tds_queue_destroy(tds_io_queue_t* queue); +tds_result_t tds_queue_synchronize(tds_io_queue_t* queue); + +tds_result_t tds_read_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, + tds_io_queue_t* queue, tds_future_t** future); +tds_result_t tds_write_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, + tds_io_queue_t* queue, tds_future_t** future); +tds_result_t tds_readv_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, tds_io_queue_t* queue, + tds_future_t** future); +tds_result_t tds_writev_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, tds_io_queue_t* queue, + tds_future_t** future); + +void tds_future_wait(tds_future_t* future); +tds_io_status_t tds_future_get_status(tds_future_t* future, int index); +void tds_future_destroy(tds_future_t* future); + +static inline tds_io_status_t tds_batch_future_get_status(tds_future_t* future, + int index) { + return tds_future_get_status(future, index); +} + +/* ========================================================================== */ +/* 4.5 Batched I/O operations */ +/* ========================================================================== */ + +tds_batch_t* tds_batch_create(void); +tds_result_t tds_batch_add(tds_batch_t* batch, + const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, OpType opt, + int* index); +tds_result_t tds_batch_addv(tds_batch_t* batch, + const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, OpType opt, int* index); +int tds_batch_get_count(const tds_batch_t* batch); +void tds_batch_destroy(tds_batch_t* batch); + +tds_result_t tds_batch_execute(tds_batch_t* batch, ssize_t* ret_arr); +tds_result_t tds_batch_submit(tds_batch_t* batch, tds_io_queue_t* queue, + tds_future_t** future); + +#ifdef __cplusplus +} +#endif + +#endif // TDSUL_H_ diff --git a/tpu_sync/tpudirect_storage/src/def_internal.h b/tpu_sync/tpudirect_storage/src/def_internal.h new file mode 100644 index 000000000..3eb6de427 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/def_internal.h @@ -0,0 +1,123 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_DEF_INTERNAL_H_ +#define TDSUL_DEF_INTERNAL_H_ + +#include +#include +#include + +#include "tdsul/def.h" + +#ifndef IOV_MAX +#define IOV_MAX 1024 +#endif + +struct tds_config { + int num_worker_threads = 4; + bool enable_io_uring = false; + bool enable_request_chunking = false; + size_t chunk_size_bytes = 0; + bool enable_p2p = false; + int worker_core_start = -1; + int worker_core_end = -1; +}; + +struct tds_storage_opts { + size_t stripe_size = 0; +}; + +struct tds_storage_handle { + std::string uri; + int fd = -1; + bool owns_fd = false; + size_t stripe_size = 0; +}; + +struct tds_buffer_handle { + int memory_type = TDS_MEM_HOST; + void* vaddr = nullptr; + int dma_buf_fd = -1; + size_t offset = 0; + size_t size = 0; + bool is_mapped = false; +}; + +struct tds_io_future { + tds_io_status_t status; +}; + +struct tds_batch { + std::vector storage_ios; + std::vector buffer_ios; + std::vector ops; +}; + +namespace tdsul { +enum class RequestType { SINGLE, BATCHED, CHUNKED }; +enum class RequestState { + // Blocked in the IO queue waiting for its turn due to FIFO ordering or batch + // barriers. + PENDING, + // Scheduled to the worker threads (actively processing or waiting in thread + // pool). + SCHEDULED, + // Processing is complete and the request is awaiting cleanup/polling by the + // user. + FINISHED +}; + +struct WorkRequest { + WorkRequest(RequestType type, tds_io_queue_t* queue, tds_future_t* future) + : type_(type), + queue_(queue), + future_(future), + state_(RequestState::PENDING) {} + virtual ~WorkRequest() = default; + + RequestType type_; + tds_io_queue_t* queue_; + tds_future_t* future_; + RequestState state_; + virtual size_t GetTransferSize() const = 0; +}; + +struct SingleWorkRequest : public WorkRequest { + SingleWorkRequest(tds_storage_io_t storage, tds_buffer_io_t buffer, OpType op, + tds_io_queue_t* queue, tds_future_t* future) + : WorkRequest(RequestType::SINGLE, queue, future), + storage_io_desc(storage), + buffer_io_desc(buffer), + op_type(op) {} + + size_t GetTransferSize() const override { + // TODO: add support for readv and writev. + if (buffer_io_desc.handle) { + return buffer_io_desc.registered.size; + } + return buffer_io_desc.raw_ptr.size; + } + + tds_storage_io_t storage_io_desc; + tds_buffer_io_t buffer_io_desc; + OpType op_type; + // TODO: we need a backpointer to the father work request in case this is + // a chunked work request. Or we use another struct for the chunked + // subrequest. +}; +// TODO: define ChunkedWorkRequest, define BatchWorkRequest. +} // namespace tdsul + +#endif // TDSUL_DEF_INTERNAL_H_ diff --git a/tpu_sync/tpudirect_storage/src/syscall_internal.h b/tpu_sync/tpudirect_storage/src/syscall_internal.h new file mode 100644 index 000000000..db9ae2c5f --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/syscall_internal.h @@ -0,0 +1,49 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_SRC_SYSCALL_INTERNAL_H_ +#define TDSUL_SRC_SYSCALL_INTERNAL_H_ + +#include +#include +#include + +#include + +namespace tdsul { + +class Syscall { + public: + virtual ~Syscall() = default; + virtual ssize_t pread(int fd, void* buf, size_t count, off_t offset) { + return ::pread(fd, buf, count, offset); + } + virtual ssize_t pwrite(int fd, const void* buf, size_t count, off_t offset) { + return ::pwrite(fd, buf, count, offset); + } + virtual ssize_t preadv(int fd, const struct iovec* iov, int iovcnt, + off_t offset) { + return ::preadv(fd, iov, iovcnt, offset); + } + virtual ssize_t pwritev(int fd, const struct iovec* iov, int iovcnt, + off_t offset) { + return ::pwritev(fd, iov, iovcnt, offset); + } +}; + +extern Syscall* g_syscall; + +} // namespace tdsul + +#endif // TDSUL_SRC_SYSCALL_INTERNAL_H_ diff --git a/tpu_sync/tpudirect_storage/src/tdsul.cpp b/tpu_sync/tpudirect_storage/src/tdsul.cpp new file mode 100644 index 000000000..65c9470c6 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/tdsul.cpp @@ -0,0 +1,841 @@ +// Copyright 2026 Google LLC. +// +// 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 "tdsul/tdsul.h" + +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +#include "absl/log/log.h" +#include "absl/strings/match.h" +#include "def_internal.h" +#include "syscall_internal.h" +#include "tdsul/def.h" +#include "util/IoConductor.h" +#include "util/IoQueue.h" +#include "util/ThreadPoolConductor.h" + +namespace tdsul { +static Syscall default_syscall; +Syscall* g_syscall = &default_syscall; +static bool g_p2p_enabled = false; +std::unique_ptr g_io_conductor; +std::mutex g_init_mutex; +bool tds_initialized = false; +bool g_use_io_uring = false; + +#ifdef TDSUL_ENABLE_IO_URING +static bool check_io_uring_support() { +#ifdef __NR_io_uring_setup + long ret = syscall(__NR_io_uring_setup, 1, nullptr); + if (ret < 0) { + if (errno == ENOSYS || errno == EPERM || errno == EACCES) { + return false; + } + // EFAULT means the syscall exists and tried to read our nullptr + return true; + } + close(ret); + return true; +#else + return false; +#endif +} +#endif + +static bool ValidateIoArgs(const char* op_name, int fd, const void* buf, + off_t file_offset, off_t buf_offset, + std::size_t io_size, std::size_t buf_size) { + if (fd < 0) { + LOG(ERROR) << op_name << " failed: invalid file descriptor: " << fd; + errno = EINVAL; + return false; + } + if (!buf) { + LOG(ERROR) << op_name << " failed: buffer pointer is null"; + errno = EINVAL; + return false; + } + if (file_offset < 0) { + LOG(ERROR) << op_name << " failed: negative file offset: " << file_offset; + errno = EINVAL; + return false; + } + if (buf_offset < 0) { + LOG(ERROR) << op_name << " failed: negative buffer offset: " << buf_offset; + errno = EINVAL; + return false; + } + if (static_cast(buf_offset) + io_size > buf_size) { + LOG(ERROR) << op_name << " failed: I/O range [" << buf_offset << ", " + << buf_offset + io_size << ") exceeds registered buffer size (" + << buf_size << ")"; + errno = EINVAL; + return false; + } + return true; +} + +static bool GetBufferInfo(const tds_buffer_io_t* bio, void** vaddr, + size_t* offset, size_t* io_size, + size_t* full_buffer_size) { + if (bio->handle) { + if (!bio->handle->vaddr) { + LOG(ERROR) << "Invalid buffer handle: vaddr is null"; + return false; + } + if (bio->registered.offset > bio->handle->size || + bio->registered.size > bio->handle->size - bio->registered.offset) { + LOG(ERROR) << "Buffer I/O out of bounds: offset=" + << bio->registered.offset << ", size=" << bio->registered.size + << ", buffer_size=" << bio->handle->size; + return false; + } + *vaddr = bio->handle->vaddr; + *offset = bio->registered.offset; + *io_size = bio->registered.size; + if (full_buffer_size) { + *full_buffer_size = bio->handle->size; + } + } else { + if (!bio->raw_ptr.vptr) { + LOG(ERROR) << "Invalid raw buffer: vptr is null"; + return false; + } + *vaddr = bio->raw_ptr.vptr; + *offset = 0; + *io_size = bio->raw_ptr.size; + if (full_buffer_size) { + *full_buffer_size = bio->raw_ptr.size; + } + } + return true; +} + +static ssize_t ReadFileToRawBuffer(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io) { + int fd = storage_io->handle->fd; + off_t file_offset = storage_io->file.file_offset; + + void* raw_buf = nullptr; + size_t buf_offset = 0; + size_t io_size = 0; + size_t full_buffer_size = 0; + if (!GetBufferInfo(buffer_io, &raw_buf, &buf_offset, &io_size, + &full_buffer_size)) { + errno = EINVAL; + return -1; + } + + if (io_size == 0) { + return 0; // No-op, return success with 0 bytes read + } + + if (!ValidateIoArgs("Read", fd, raw_buf, file_offset, buf_offset, io_size, + full_buffer_size)) { + return -1; + } + + char* buf_ptr = static_cast(raw_buf) + buf_offset; + + std::size_t total_to_read = io_size; + + ssize_t bytes_read; + do { + bytes_read = g_syscall->pread(fd, buf_ptr, total_to_read, file_offset); + } while (bytes_read < 0 && errno == EINTR); + + if (bytes_read < 0) { + const int saved_errno = errno; + LOG(ERROR) << "pread failed: " << std::strerror(saved_errno); + errno = saved_errno; + return -1; + } + + return bytes_read; +} + +static ssize_t WriteFileFromRawBuffer(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io) { + int fd = storage_io->handle->fd; + off_t file_offset = storage_io->file.file_offset; + + void* raw_buf = nullptr; + size_t buf_offset = 0; + size_t io_size = 0; + size_t full_buffer_size = 0; + if (!GetBufferInfo(buffer_io, &raw_buf, &buf_offset, &io_size, + &full_buffer_size)) { + errno = EINVAL; + return -1; + } + + if (io_size == 0) { + return 0; // No-op, return success with 0 bytes written + } + + if (!ValidateIoArgs("Write", fd, raw_buf, file_offset, buf_offset, io_size, + full_buffer_size)) { + return -1; + } + + char* buf_ptr = static_cast(raw_buf) + buf_offset; + + std::size_t total_to_write = io_size; + + ssize_t bytes_written; + do { + bytes_written = g_syscall->pwrite(fd, buf_ptr, total_to_write, file_offset); + } while (bytes_written < 0 && errno == EINTR); + + if (bytes_written < 0) { + const int saved_errno = errno; + LOG(ERROR) << "pwrite failed: " << std::strerror(saved_errno); + errno = saved_errno; + return -1; + } + + return bytes_written; +} + +static ssize_t ReadvFileToRawBuffers(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_iov, + int iovcnt) { + if (iovcnt > IOV_MAX) { + LOG(ERROR) << "Readv failed: iovcnt (" << iovcnt << ") exceeds IOV_MAX (" + << IOV_MAX << ")"; + return -1; + } + + int fd = storage_io->handle->fd; + off_t file_offset = storage_io->file.file_offset; + + std::vector iov(iovcnt); + std::size_t total_to_read = 0; + + for (int i = 0; i < iovcnt; ++i) { + const tds_buffer_io_t& bio = buffer_iov[i]; + void* raw_buf = nullptr; + std::size_t io_size = 0; + std::size_t buf_offset = 0; + std::size_t full_buffer_size = 0; + if (!GetBufferInfo(&bio, &raw_buf, &buf_offset, &io_size, + &full_buffer_size)) { + errno = EINVAL; + return -1; + } + + if (!ValidateIoArgs("Readv", fd, raw_buf, file_offset, buf_offset, io_size, + full_buffer_size)) { + return -1; + } + + char* buf_ptr = static_cast(raw_buf) + buf_offset; + iov[i].iov_base = buf_ptr; + iov[i].iov_len = io_size; + + if (io_size > SSIZE_MAX - total_to_read) { + LOG(ERROR) << "Readv failed: total size exceeds SSIZE_MAX"; + errno = EINVAL; + return -1; + } + total_to_read += io_size; + } + + if (total_to_read == 0) { + return 0; + } + + ssize_t bytes_read; + do { + bytes_read = g_syscall->preadv(fd, iov.data(), iov.size(), file_offset); + } while (bytes_read < 0 && errno == EINTR); + + if (bytes_read < 0) { + const int saved_errno = errno; + LOG(ERROR) << "preadv failed: " << std::strerror(saved_errno); + errno = saved_errno; + return -1; + } + + return bytes_read; +} + +static ssize_t WritevFileFromRawBuffers(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_iov, + int iovcnt) { + if (iovcnt > IOV_MAX) { + LOG(ERROR) << "Writev failed: iovcnt (" << iovcnt << ") exceeds IOV_MAX (" + << IOV_MAX << ")"; + return -1; + } + + int fd = storage_io->handle->fd; + off_t file_offset = storage_io->file.file_offset; + + std::vector iov(iovcnt); + std::size_t total_to_write = 0; + + for (int i = 0; i < iovcnt; ++i) { + const tds_buffer_io_t& bio = buffer_iov[i]; + void* raw_buf = nullptr; + std::size_t io_size = 0; + std::size_t buf_offset = 0; + std::size_t full_buffer_size = 0; + if (!GetBufferInfo(&bio, &raw_buf, &buf_offset, &io_size, + &full_buffer_size)) { + errno = EINVAL; + return -1; + } + + if (!ValidateIoArgs("Writev", fd, raw_buf, file_offset, buf_offset, io_size, + full_buffer_size)) { + return -1; + } + + char* buf_ptr = static_cast(raw_buf) + buf_offset; + iov[i].iov_base = buf_ptr; + iov[i].iov_len = io_size; + + if (io_size > SSIZE_MAX - total_to_write) { + LOG(ERROR) << "Writev failed: total size exceeds SSIZE_MAX"; + errno = EINVAL; + return -1; + } + total_to_write += io_size; + } + + if (total_to_write == 0) { + return 0; + } + + ssize_t bytes_written; + do { + bytes_written = g_syscall->pwritev(fd, iov.data(), iov.size(), file_offset); + } while (bytes_written < 0 && errno == EINTR); + + if (bytes_written < 0) { + const int saved_errno = errno; + LOG(ERROR) << "pwritev failed: " << std::strerror(saved_errno); + errno = saved_errno; + return -1; + } + + return bytes_written; +} +} // namespace tdsul + +/* ========================================================================== */ +/* 4.1 Initialization & Lifecycle Management */ +/* ========================================================================== */ + +tds_config_t* tds_config_create(void) { + return new (std::nothrow) tds_config_t(); +} + +void tds_config_destroy(tds_config_t* config) { delete config; } + +tds_result_t tds_config_set_int(tds_config_t* config, const char* key, + int value) { + if (!config || !key) return TDS_ERROR_INVALID_PARAMETER; + if (std::strcmp(key, "num_worker_threads") == 0) { + config->num_worker_threads = value; + return TDS_SUCCESS; + } + if (std::strcmp(key, "enable_io_uring") == 0) { + config->enable_io_uring = (value != 0); + return TDS_SUCCESS; + } + if (std::strcmp(key, "worker_core_start") == 0) { + config->worker_core_start = value; + return TDS_SUCCESS; + } + if (std::strcmp(key, "worker_core_end") == 0) { + config->worker_core_end = value; + return TDS_SUCCESS; + } + // TODO: Check if the hardware is physically present and the kernel driver is + // loaded to avoid relying solely on the Python layer. + if (std::strcmp(key, "enable_p2p") == 0) { + config->enable_p2p = (value != 0); + return TDS_SUCCESS; + } + return TDS_ERROR_INVALID_PARAMETER; +} + +tds_result_t tds_config_set_string(tds_config_t* config, const char* key, + const char* value) { + if (!config || !key || !value) return TDS_ERROR_INVALID_PARAMETER; + return TDS_SUCCESS; +} + +tds_storage_opts_t* tds_storage_opts_create(void) { + return new (std::nothrow) tds_storage_opts_t(); +} + +void tds_storage_opts_destroy(tds_storage_opts_t* opts) { delete opts; } + +tds_result_t tds_storage_opts_set_int(tds_storage_opts_t* opts, const char* key, + int value) { + if (!opts || !key) return TDS_ERROR_INVALID_PARAMETER; + if (std::strcmp(key, "stripe_size") == 0) { + opts->stripe_size = static_cast(value); + return TDS_SUCCESS; + } + return TDS_ERROR_INVALID_PARAMETER; +} + +tds_result_t tds_storage_opts_set_string(tds_storage_opts_t* opts, + const char* key, const char* value) { + if (!opts || !key || !value) return TDS_ERROR_INVALID_PARAMETER; + return TDS_SUCCESS; +} + +tds_result_t tds_init(const tds_config_t* config) { + std::lock_guard lock(tdsul::g_init_mutex); + if (tdsul::tds_initialized) { + LOG(WARNING) << "TDS is already initialized."; + return TDS_SUCCESS; + } + + int num_threads = 4; + int core_start = -1; + int core_end = -1; + tdsul::g_use_io_uring = false; + tdsul::g_p2p_enabled = false; + bool enable_io_uring = false; + + if (config) { + tdsul::g_p2p_enabled = config->enable_p2p; + if (config->num_worker_threads > 0) { + num_threads = config->num_worker_threads; + } + core_start = config->worker_core_start; + core_end = config->worker_core_end; + enable_io_uring = config->enable_io_uring; + } + +#ifdef TDSUL_ENABLE_IO_URING + if (enable_io_uring) { + if (tdsul::check_io_uring_support()) { + tdsul::g_use_io_uring = true; + LOG(INFO) << "io_uring is supported and enabled by user config."; + } else { + LOG(WARNING) << "io_uring is requested but not supported by the system. " + "Falling back to worker threads."; + } + } else { + LOG(INFO) << "io_uring is not enabled. Using worker threads."; + } + + if (tdsul::g_use_io_uring) { + tdsul::g_io_conductor = std::make_unique(); + } else { + tdsul::g_io_conductor = std::make_unique( + num_threads, core_start, core_end); + } +#else + if (enable_io_uring) { + LOG(WARNING) << "io_uring requested but disabled at compile time! " + "Falling back to worker threads."; + } else { + LOG(INFO) << "io_uring is not enabled. Using worker threads."; + } + + tdsul::g_io_conductor = std::make_unique( + num_threads, core_start, core_end); +#endif + + tdsul::g_io_conductor->Start(); + tdsul::tds_initialized = true; + return TDS_SUCCESS; +} + +tds_result_t tds_shutdown(void) { + std::lock_guard lock(tdsul::g_init_mutex); + if (tdsul::g_io_conductor) { + tdsul::g_io_conductor->Stop(); + tdsul::g_io_conductor.reset(); + } + tdsul::g_use_io_uring = false; + tdsul::tds_initialized = false; + return TDS_SUCCESS; +} + +/* ========================================================================== */ +/* 4.2 Registration Management */ +/* ========================================================================== */ + +tds_result_t tds_buffer_register_vaddr(void* vaddr, int memory_type, + size_t size, + tds_buffer_handle_t** buf_handle) { + if (!buf_handle) { + LOG(ERROR) << "tds_buffer_register_vaddr failed: buf_handle output pointer " + "is null"; + return TDS_ERROR_INVALID_PARAMETER; + } + + if (vaddr == nullptr) { + LOG(ERROR) << "tds_buffer_register_vaddr failed: vaddr is null"; + return TDS_ERROR_INVALID_PARAMETER; + } + + if (memory_type == TDS_MEM_DEVICE && !tdsul::g_p2p_enabled) { + LOG(ERROR) << "tds_buffer_register_vaddr failed: P2P device memory " + "registration is not supported in this environment."; + *buf_handle = nullptr; + return TDS_ERROR_P2P_UNSUPPORTED; + } + + auto* handle = new tds_buffer_handle_t(); + handle->memory_type = memory_type; + handle->vaddr = vaddr; + handle->size = size; + handle->is_mapped = false; + + *buf_handle = handle; + LOG(INFO) << "Successfully registered raw DRAM host buffer at " << vaddr + << " (size: " << size << ")"; + return TDS_SUCCESS; +} + +tds_result_t tds_buffer_register_dmabuf(int dma_buf_fd, size_t offset, + int memory_type, size_t size, + tds_buffer_handle_t** buf_handle) { + if (!buf_handle) { + LOG(ERROR) << "tds_buffer_register_dmabuf failed: buf_handle output " + "pointer is null"; + return TDS_ERROR_INVALID_PARAMETER; + } + + if (dma_buf_fd < 0) { + LOG(ERROR) << "tds_buffer_register_dmabuf failed: dma_buf_fd is invalid: " + << dma_buf_fd; + return TDS_ERROR_INVALID_PARAMETER; + } + + if (memory_type == TDS_MEM_DEVICE && !tdsul::g_p2p_enabled) { + LOG(ERROR) << "tds_buffer_register_dmabuf failed: P2P device memory " + "registration is not supported in this environment."; + *buf_handle = nullptr; + return TDS_ERROR_P2P_UNSUPPORTED; + } + + void* mapped_addr = mmap(nullptr, size, PROT_READ | PROT_WRITE, MAP_SHARED, + dma_buf_fd, static_cast(offset)); + if (mapped_addr == MAP_FAILED) { + LOG(ERROR) + << "tds_buffer_register_dmabuf failed: mmap failed for dma_buf_fd: " + << dma_buf_fd << " at offset: " << offset << " (size: " << size + << "): " << std::strerror(errno); + return TDS_ERROR_IO; + } + + auto* handle = new tds_buffer_handle_t(); + handle->memory_type = TDS_MEM_HOST; + handle->vaddr = mapped_addr; + handle->dma_buf_fd = dma_buf_fd; + handle->offset = offset; + handle->size = size; + handle->is_mapped = true; + + *buf_handle = handle; + LOG(INFO) << "Successfully registered mapped DRAM dmabuf at " << mapped_addr + << " (size: " << size << ", dma_buf_fd: " << dma_buf_fd << ")"; + return TDS_SUCCESS; +} + +tds_result_t tds_buffer_deregister(tds_buffer_handle_t* buf_handle) { + if (!buf_handle) { + LOG(ERROR) << "tds_buffer_deregister failed: buf_handle is null"; + return TDS_ERROR_INVALID_PARAMETER; + } + + if (buf_handle->is_mapped && buf_handle->vaddr != nullptr) { + if (munmap(buf_handle->vaddr, buf_handle->size) != 0) { + LOG(ERROR) << "munmap failed for address " << buf_handle->vaddr << ": " + << std::strerror(errno); + } + } + + delete buf_handle; + LOG(INFO) << "Successfully deregistered buffer handle"; + return TDS_SUCCESS; +} + +tds_result_t tds_storage_handle_register( + const tds_storage_descr_t* descr, tds_storage_handle_t** storage_handle) { + if (!descr || !descr->uri || !storage_handle) + return TDS_ERROR_INVALID_PARAMETER; + + std::string uri(descr->uri); + int fd = -1; + bool owns_fd = false; + + if (uri.rfind("fd://", 0) == 0) { + try { + fd = std::stoi(uri.substr(5)); + } catch (...) { + LOG(ERROR) << "Invalid file descriptor in URI: " << uri; + return TDS_ERROR_INVALID_PARAMETER; + } + owns_fd = false; + } else { + std::string path; + if (uri.rfind("file://", 0) == 0) { + path = uri.substr(7); + } else if (absl::StrContains(uri, "://")) { + LOG(WARNING) << "Storage URI scheme is not supported: " << uri; + return TDS_ERROR_UNSUPPORTED; + } else { + path = uri; + } + + fd = ::open(path.c_str(), O_RDWR | O_CREAT, 0644); + if (fd < 0) { + fd = ::open(path.c_str(), O_RDONLY); + } + if (fd < 0) { + LOG(ERROR) << "Failed to open file path '" << path + << "': " << std::strerror(errno); + return TDS_ERROR_IO; + } + owns_fd = true; + } + + auto* handle = new tds_storage_handle_t(); + handle->uri = uri; + handle->fd = fd; + handle->owns_fd = owns_fd; + if (descr->options) { + handle->stripe_size = descr->options->stripe_size; + } + + *storage_handle = handle; + LOG(INFO) << "Successfully registered storage handle for URI: " << uri + << " (fd: " << fd << ", owns_fd: " << owns_fd << ")"; + return TDS_SUCCESS; +} + +tds_result_t tds_storage_handle_deregister( + tds_storage_handle_t* storage_handle) { + if (!storage_handle) return TDS_ERROR_INVALID_PARAMETER; + if (storage_handle->owns_fd && storage_handle->fd >= 0) { + if (::close(storage_handle->fd) != 0) { + LOG(ERROR) << "Failed to close storage handle fd " << storage_handle->fd + << ": " << std::strerror(errno); + } + } + delete storage_handle; + LOG(INFO) << "Successfully deregistered storage handle"; + return TDS_SUCCESS; +} + +/* ========================================================================== */ +/* 4.3 Synchronous I/O */ +/* ========================================================================== */ + +ssize_t tds_read(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io) { + if (!storage_io || !storage_io->handle || !buffer_io) return -1; + return tdsul::ReadFileToRawBuffer(storage_io, buffer_io); +} + +ssize_t tds_write(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io) { + if (!storage_io || !storage_io->handle || !buffer_io) return -1; + return tdsul::WriteFileFromRawBuffer(storage_io, buffer_io); +} + +ssize_t tds_readv(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, int num_buffers) { + if (!storage_io || !storage_io->handle || !buffer_io_arr || num_buffers <= 0) + return -1; + return tdsul::ReadvFileToRawBuffers(storage_io, buffer_io_arr, num_buffers); +} + +ssize_t tds_writev(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, int num_buffers) { + if (!storage_io || !storage_io->handle || !buffer_io_arr || num_buffers <= 0) + return -1; + return tdsul::WritevFileFromRawBuffers(storage_io, buffer_io_arr, + num_buffers); +} + +/* ========================================================================== */ +/* 4.4 Asynchronous I/O & Queue Management */ +/* ========================================================================== */ + +tds_result_t tds_queue_create(tds_io_queue_t** queue) { + if (!queue) return TDS_ERROR_INVALID_PARAMETER; + *queue = nullptr; + return TDS_ERROR_UNSUPPORTED; +} + +tds_result_t tds_queue_destroy(tds_io_queue_t* queue) { + if (!queue) return TDS_ERROR_INVALID_PARAMETER; + delete reinterpret_cast(queue); + return TDS_SUCCESS; +} + +tds_result_t tds_queue_synchronize(tds_io_queue_t* queue) { + if (!queue) return TDS_ERROR_INVALID_PARAMETER; + return TDS_ERROR_UNSUPPORTED; +} + +tds_result_t tds_read_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, + tds_io_queue_t* queue, tds_future_t** future) { + if (future) *future = nullptr; + if (!storage_io || !buffer_io) return TDS_ERROR_INVALID_PARAMETER; + + auto* new_future = new tds_future_t(); + new_future->status.status = TDS_PENDING; + + auto* req = new tdsul::SingleWorkRequest(*storage_io, *buffer_io, TDS_OP_READ, + queue, new_future); + + if (queue) { + auto* io_queue = reinterpret_cast(queue); + io_queue->Push(req); + } else { + if (tdsul::g_io_conductor) { + tdsul::g_io_conductor->Dispatch(req); + } + } + + if (future) *future = new_future; + return TDS_SUCCESS; +} + +tds_result_t tds_write_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, + tds_io_queue_t* queue, tds_future_t** future) { + if (future) *future = nullptr; + if (!storage_io || !buffer_io) return TDS_ERROR_INVALID_PARAMETER; + + auto* new_future = new tds_future_t(); + new_future->status.status = TDS_PENDING; + + auto* req = new tdsul::SingleWorkRequest(*storage_io, *buffer_io, + TDS_OP_WRITE, queue, new_future); + + if (queue) { + auto* io_queue = reinterpret_cast(queue); + io_queue->Push(req); + } else { + if (tdsul::g_io_conductor) { + tdsul::g_io_conductor->Dispatch(req); + } + } + + if (future) *future = new_future; + return TDS_SUCCESS; +} + +tds_result_t tds_readv_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, tds_io_queue_t* queue, + tds_future_t** future) { + if (future) *future = nullptr; + if (!queue || !storage_io || !buffer_io_arr || num_buffers <= 0) + return TDS_ERROR_INVALID_PARAMETER; + return TDS_ERROR_UNSUPPORTED; +} + +tds_result_t tds_writev_async(const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, tds_io_queue_t* queue, + tds_future_t** future) { + if (future) *future = nullptr; + if (!queue || !storage_io || !buffer_io_arr || num_buffers <= 0) + return TDS_ERROR_INVALID_PARAMETER; + return TDS_ERROR_UNSUPPORTED; +} + +void tds_future_wait(tds_future_t* future) { (void)future; } + +tds_io_status_t tds_future_get_status(tds_future_t* future, int index) { + (void)index; + if (!future) { + return {TDS_ERROR_INVALID_PARAMETER, EINVAL, 0}; + } + return {TDS_ERROR_UNSUPPORTED, 0, 0}; +} + +void tds_future_destroy(tds_future_t* future) { delete future; } + +/* ========================================================================== */ +/* 4.5 Batched I/O operations */ +/* ========================================================================== */ + +tds_batch_t* tds_batch_create(void) { return nullptr; } + +tds_result_t tds_batch_add(tds_batch_t* batch, + const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io, OpType opt, + int* index) { + (void)batch; + (void)storage_io; + (void)buffer_io; + (void)opt; + if (index) *index = -1; + return TDS_ERROR_UNSUPPORTED; +} + +tds_result_t tds_batch_addv(tds_batch_t* batch, + const tds_storage_io_t* storage_io, + const tds_buffer_io_t* buffer_io_arr, + int num_buffers, OpType opt, int* index) { + (void)batch; + (void)storage_io; + (void)buffer_io_arr; + (void)num_buffers; + (void)opt; + if (index) *index = -1; + return TDS_ERROR_UNSUPPORTED; +} + +int tds_batch_get_count(const tds_batch_t* batch) { + (void)batch; + return 0; +} + +void tds_batch_destroy(tds_batch_t* batch) { (void)batch; } + +tds_result_t tds_batch_execute(tds_batch_t* batch, ssize_t* ret_arr) { + (void)batch; + (void)ret_arr; + return TDS_ERROR_UNSUPPORTED; +} + +tds_result_t tds_batch_submit(tds_batch_t* batch, tds_io_queue_t* queue, + tds_future_t** future) { + (void)batch; + (void)queue; + if (future) *future = nullptr; + return TDS_ERROR_UNSUPPORTED; +} diff --git a/tpu_sync/tpudirect_storage/src/util/IoConductor.h b/tpu_sync/tpudirect_storage/src/util/IoConductor.h new file mode 100644 index 000000000..d3b3018c7 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/IoConductor.h @@ -0,0 +1,35 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_SRC_UTIL_IO_CONDUCTOR_H_ +#define TDSUL_SRC_UTIL_IO_CONDUCTOR_H_ + +#include "def_internal.h" + +namespace tdsul { + +enum class IoConductorType { THREAD_POOL, IO_URING }; + +class IoConductor { + public: + virtual ~IoConductor() = default; + virtual void Start() = 0; + virtual void Stop() = 0; + virtual void Dispatch(WorkRequest* request) = 0; + virtual IoConductorType GetType() const = 0; +}; + +} // namespace tdsul + +#endif // TDSUL_SRC_UTIL_IO_CONDUCTOR_H_ diff --git a/tpu_sync/tpudirect_storage/src/util/IoQueue.cpp b/tpu_sync/tpudirect_storage/src/util/IoQueue.cpp new file mode 100644 index 000000000..e0147186d --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/IoQueue.cpp @@ -0,0 +1,98 @@ +// Copyright 2026 Google LLC. +// +// 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 "IoQueue.h" + +#include +#include + +#include "def_internal.h" + +namespace tdsul { + +IoQueue::IoQueue() = default; + +IoQueue::~IoQueue() = default; + +void IoQueue::Push(WorkRequest* request) { + std::lock_guard lock(mutex_); + request->state_ = RequestState::PENDING; + total_transfer_size_ += request->GetTransferSize(); + queue_.push_back(request); +} + +bool IoQueue::Pop(WorkRequest*& request) { + std::lock_guard lock(mutex_); + if (!queue_.empty() && queue_.front()->state_ == RequestState::FINISHED) { + request = queue_.front(); + total_transfer_size_ -= request->GetTransferSize(); + queue_.pop_front(); + return true; + } + return false; +} + +bool IoQueue::Front(WorkRequest*& request) const { + std::lock_guard lock(mutex_); + if (queue_.empty()) { + return false; + } + request = queue_.front(); + return true; +} + +bool IoQueue::Empty() const { + std::lock_guard lock(mutex_); + return queue_.empty(); +} + +std::size_t IoQueue::Size() const { + std::lock_guard lock(mutex_); + return queue_.size(); +} + +bool IoQueue::TestAndSetAssigned() { + bool expected = false; + return is_assigned_.compare_exchange_strong(expected, true); +} + +WorkRequest* IoQueue::CheckAndGetNextSchedulableRequest() { + std::lock_guard lock(mutex_); + + if (queue_.empty()) { + is_assigned_.store(false); + return nullptr; + } + + WorkRequest* req = queue_.front(); + if (req->state_ == RequestState::PENDING) { + req->state_ = RequestState::SCHEDULED; + return req; + } + + is_assigned_.store(false); + return nullptr; +} + +void IoQueue::SetWorker(class WorkerThread* worker) { + std::lock_guard lock(mutex_); + assigned_worker_ = worker; +} + +size_t IoQueue::GetTotalTransferSize() const { + std::lock_guard lock(mutex_); + return total_transfer_size_; +} + +} // namespace tdsul diff --git a/tpu_sync/tpudirect_storage/src/util/IoQueue.h b/tpu_sync/tpudirect_storage/src/util/IoQueue.h new file mode 100644 index 000000000..c493a859d --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/IoQueue.h @@ -0,0 +1,98 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_SRC_UTIL_IO_QUEUE_H_ +#define TDSUL_SRC_UTIL_IO_QUEUE_H_ + +#include +#include +#include +#include + +#include "../def_internal.h" + +namespace tdsul { + +/** + * @brief IoQueue manages pending and scheduled WorkRequests, ensuring FIFO + * order is obeyed given a mixed workload of single and batch requests. + * + * Batch requests are required to obey strict FIFO ordering while still + * utilizing the worker thread pool as much as possible. Therefore, a batch + * request, as well as any single request that comes after a batch, must wait + * for a clear signal that the previous work requests have finished. Without + * batch requests, single requests could be scheduled to a worker thread's local + * queue immediately. But batch requests make it complex. + */ +class IoQueue { + public: + IoQueue(); + ~IoQueue(); + + // Disallow copy and move semantics for safety and simplicity + IoQueue(const IoQueue&) = delete; + IoQueue& operator=(const IoQueue&) = delete; + IoQueue(IoQueue&&) = delete; + IoQueue& operator=(IoQueue&&) = delete; + + // Returns true if the queue transitions from unassigned to assigned. + bool TestAndSetAssigned(); + + // Used by the worker thread to get the next request without holding a lock + // long. + WorkRequest* CheckAndGetNextSchedulableRequest(); + + /** + * @brief Enqueues a work request into the queue in a thread-safe manner. + */ + void Push(WorkRequest* request); + + /** + * @brief Dequeues a work request from the queue in a thread-safe manner. + * @return true if a request was popped, false if queue was empty. + */ + bool Pop(WorkRequest*& request); + + /** + * @brief Retrieves the work request at the front of the queue without + * removing it. + * @return true if a request was retrieved, false if queue was empty. + */ + bool Front(WorkRequest*& request) const; + + /** + * @brief Checks if the queue is empty in a thread-safe manner. + */ + bool Empty() const; + + /** + * @brief Returns the number of requests in the queue in a thread-safe manner. + */ + std::size_t Size() const; + + void SetWorker(class WorkerThread* worker); + size_t GetTotalTransferSize() const; + + private: + // Mutex to guard internal queue state for thread-safety. + mutable std::mutex mutex_; + std::deque queue_; + std::atomic is_assigned_{false}; + class WorkerThread* assigned_worker_ = nullptr; + size_t total_transfer_size_ = 0; +}; + +} // namespace tdsul + +#endif // TDSUL_SRC_UTIL_IO_QUEUE_H_ diff --git a/tpu_sync/tpudirect_storage/src/util/IoUringConductor.cpp b/tpu_sync/tpudirect_storage/src/util/IoUringConductor.cpp new file mode 100644 index 000000000..92da432cb --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/IoUringConductor.cpp @@ -0,0 +1,44 @@ +// Copyright 2026 Google LLC. +// +// 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 "util/IoUringConductor.h" + +#ifdef TDSUL_ENABLE_IO_URING + +#include "absl/log/log.h" +#include "def_internal.h" +#include "util/IoQueue.h" + +namespace tdsul { + +IoUringConductor::IoUringConductor() { + LOG(INFO) << "IoUringConductor created."; +} + +IoUringConductor::~IoUringConductor() { Stop(); } + +void IoUringConductor::Start() { + LOG(INFO) << "IoUringConductor Start (Stub)."; +} + +void IoUringConductor::Stop() { LOG(INFO) << "IoUringConductor Stop (Stub)."; } + +void IoUringConductor::Dispatch(WorkRequest* request) { + // will implement it in later commit. + delete request; +} + +} // namespace tdsul + +#endif // TDSUL_ENABLE_IO_URING diff --git a/tpu_sync/tpudirect_storage/src/util/IoUringConductor.h b/tpu_sync/tpudirect_storage/src/util/IoUringConductor.h new file mode 100644 index 000000000..226a12787 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/IoUringConductor.h @@ -0,0 +1,39 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_SRC_UTIL_IO_URING_CONDUCTOR_H_ +#define TDSUL_SRC_UTIL_IO_URING_CONDUCTOR_H_ + +#include "util/IoConductor.h" + +#ifdef TDSUL_ENABLE_IO_URING + +namespace tdsul { + +class IoUringConductor : public IoConductor { + public: + IoUringConductor(); + ~IoUringConductor() override; + + void Start() override; + void Stop() override; + void Dispatch(WorkRequest* request) override; + IoConductorType GetType() const override { return IoConductorType::IO_URING; } +}; + +} // namespace tdsul + +#endif // TDSUL_ENABLE_IO_URING + +#endif // TDSUL_SRC_UTIL_IO_URING_CONDUCTOR_H_ diff --git a/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.cpp b/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.cpp new file mode 100644 index 000000000..6be2e5596 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.cpp @@ -0,0 +1,240 @@ +// Copyright 2026 Google LLC. +// +// 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 "util/ThreadPoolConductor.h" + +#include +#include + +#include +#include +#include +#include +#include + +#include "absl/log/log.h" +#include "def_internal.h" +#include "tdsul/def.h" +#include "util/IoConductor.h" + +namespace tdsul { + +extern std::unique_ptr g_io_conductor; + +WorkerThread::WorkerThread(int id) + : id_(id), thread_valid_(false), stop_flag_(false) {} + +WorkerThread::~WorkerThread() { Stop(); } + +void* WorkerThread::ThreadEntry(void* arg) { + static_cast(arg)->Run(); + return nullptr; +} + +void WorkerThread::Start(int core_id) { + pthread_attr_t attr; + pthread_attr_init(&attr); + + if (core_id >= 0) { + cpu_set_t cpuset; + CPU_ZERO(&cpuset); + CPU_SET(core_id, &cpuset); + pthread_attr_setaffinity_np(&attr, sizeof(cpu_set_t), &cpuset); + } + + int rc = pthread_create(&thread_, &attr, ThreadEntry, this); + pthread_attr_destroy(&attr); + + if (rc != 0 && core_id >= 0) { + // `pthread_create` fails with EINVAL when the requested CPU lies outside + // the process's allowed cpuset. That happens routinely under restricted + // cpusets (Forge test sandboxes, Borg jobs with a CPU mask), because + // `sysconf(_SC_NPROCESSORS_ONLN)` reports the machine's cores rather than + // the cores this process may actually run on. + // + // Falling back to an unpinned thread keeps the pool functional. Without + // this, no worker starts at all and the pool silently accepts work it will + // never execute. + LOG(WARNING) << "pthread_create with affinity to core " << core_id + << " failed for worker " << id_ << " (" << rc + << "); retrying without CPU affinity"; + pthread_attr_t unpinned_attr; + pthread_attr_init(&unpinned_attr); + rc = pthread_create(&thread_, &unpinned_attr, ThreadEntry, this); + pthread_attr_destroy(&unpinned_attr); + } + + if (rc != 0) { + LOG(ERROR) << "Error calling pthread_create for worker " << id_ << ": " + << rc; + } else { + thread_valid_ = true; + if (core_id >= 0) { + LOG(INFO) << "Set worker " << id_ << " affinity to core " << core_id; + } + } +} + +void WorkerThread::Stop() { + { + std::lock_guard lock(mutex_); + stop_flag_ = true; + } + cv_.notify_one(); + if (thread_valid_) { + pthread_join(thread_, nullptr); + thread_valid_ = false; + } + + // `Run()` drains the queue before exiting, so normally nothing is left here. + // Requests do remain when the worker thread never started (see the affinity + // fallback in `Start()`); without this the pool would accept work, never + // execute it, and leak every queued `WorkRequest`. + // + // Ownership mirrors `Run()`: requests owned by an `IoQueue` are freed by that + // queue, so only unowned requests are deleted here. + std::queue abandoned; + { + std::lock_guard lock(mutex_); + abandoned.swap(task_queue_); + } + while (!abandoned.empty()) { + WorkRequest* request = abandoned.front(); + abandoned.pop(); + if (request == nullptr) continue; + if (request->future_) { + request->future_->status.status = TDS_ERROR_IO; + } + if (!request->queue_) { + delete request; + } + } +} + +void WorkerThread::EnqueueTask(WorkRequest* request) { + { + std::lock_guard lock(mutex_); + task_queue_.push(request); + } + cv_.notify_one(); +} + +size_t WorkerThread::GetPendingTaskCount() const { + std::lock_guard lock(mutex_); + return task_queue_.size(); +} + +void WorkerThread::Run() { + while (true) { + WorkRequest* request = nullptr; + { + std::unique_lock lock(mutex_); + cv_.wait(lock, [this] { return stop_flag_ || !task_queue_.empty(); }); + + if (stop_flag_ && task_queue_.empty()) { + break; + } + + request = task_queue_.front(); + task_queue_.pop(); + } + + if (request->type_ == RequestType::SINGLE) { + // TODO: Process the IO request synchronously here + VLOG(3) << "Worker " << std::to_string(id_) << " processing request."; + if (request->future_) { + request->future_->status.status = TDS_SUCCESS; + } + } + if (!request->queue_) { + // Free the request after processing only if it's not managed by a queue + delete request; + } + } +} + +ThreadPoolConductor::ThreadPoolConductor(int num_threads, int core_start, + int core_end) + : num_threads_(num_threads), core_start_(core_start), core_end_(core_end) { + for (int i = 0; i < num_threads_; ++i) { + workers_.push_back(std::make_unique(i)); + } +} + +ThreadPoolConductor::~ThreadPoolConductor() { Stop(); } + +void ThreadPoolConductor::Start() { + int num_cores = sysconf(_SC_NPROCESSORS_ONLN); + if (num_cores <= 0) { + num_cores = 1; + } + + int actual_core_start = + (core_start_ >= 0 && core_start_ < num_cores) ? core_start_ : 0; + int actual_core_end = (core_end_ >= 0 && core_end_ < num_cores && + core_end_ >= actual_core_start) + ? core_end_ + : num_cores - 1; + int core_range = actual_core_end - actual_core_start + 1; + + if (core_start_ >= 0 && core_end_ >= 0 && core_range != num_threads_) { + LOG(FATAL) << "ThreadPool size (" << num_threads_ + << ") does not match configured core range size (" << core_range + << ")."; + } + + // Only pin when the caller explicitly configured a core range via + // `tds_init`'s "worker_core_start"/"worker_core_end". Pinning by default is + // actively harmful: the core indices are derived from + // `sysconf(_SC_NPROCESSORS_ONLN)` (machine-wide), which says nothing about + // which CPUs this process is permitted to run on under a restricted cpuset. + const bool pin_to_cores = (core_start_ >= 0 && core_end_ >= 0); + + for (int i = 0; i < num_threads_; ++i) { + // TODO: Set thread affinity according to the PCIe topology, which should be + // attached to the closest core to NIC or SSD. For now, simple round-robin + // assignment within the specified core range + workers_[i]->Start(pin_to_cores ? actual_core_start + (i % core_range) + : -1); + } +} + +void ThreadPoolConductor::Stop() { + for (auto& worker : workers_) { + worker->Stop(); + } +} + +void ThreadPoolConductor::Dispatch(WorkRequest* request) { + if (workers_.empty()) { + LOG(ERROR) << "No worker threads available in ThreadPool."; + // Clean up request if we can't process it to avoid leaks + delete request; + return; + } + size_t min_pending = static_cast(-1); + int best_worker_idx = 0; + + for (int i = 0; i < num_threads_; ++i) { + size_t pending = workers_[i]->GetPendingTaskCount(); + if (pending < min_pending) { + min_pending = pending; + best_worker_idx = i; + } + } + + workers_[best_worker_idx]->EnqueueTask(request); +} + +} // namespace tdsul diff --git a/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.h b/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.h new file mode 100644 index 000000000..7a110fb02 --- /dev/null +++ b/tpu_sync/tpudirect_storage/src/util/ThreadPoolConductor.h @@ -0,0 +1,91 @@ +// Copyright 2026 Google LLC. +// +// 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. + +#ifndef TDSUL_SRC_UTIL_THREAD_POOL_CONDUCTOR_H_ +#define TDSUL_SRC_UTIL_THREAD_POOL_CONDUCTOR_H_ + +#include + +#include +#include +#include +#include +#include +#include + +#include "util/IoConductor.h" +#include "util/IoQueue.h" + +namespace tdsul { + +class WorkerThread { + public: + WorkerThread(int id); + ~WorkerThread(); + + // Disallow copy and move semantics + WorkerThread(const WorkerThread&) = delete; + WorkerThread& operator=(const WorkerThread&) = delete; + + void Start(int core_id = -1); + void Stop(); + + void EnqueueTask(WorkRequest* request); + size_t GetPendingTaskCount() const; + + private: + void Run(); + static void* ThreadEntry(void* arg); + + int id_; + std::queue task_queue_; + mutable std::mutex mutex_; + std::condition_variable cv_; + pthread_t thread_; + bool thread_valid_; + std::atomic stop_flag_; +}; + +class ThreadPoolConductor : public IoConductor { + public: + explicit ThreadPoolConductor(int num_threads, int core_start = -1, + int core_end = -1); + ~ThreadPoolConductor(); + + // Disallow copy and move semantics + ThreadPoolConductor(const ThreadPoolConductor&) = delete; + ThreadPoolConductor& operator=(const ThreadPoolConductor&) = delete; + + void Start() override; + void Stop() override; + + void Dispatch(WorkRequest* request) override; + void AssignQueue(IoQueue* queue); + IoConductorType GetType() const override { + return IoConductorType::THREAD_POOL; + } + + private: + std::vector> workers_; + int num_threads_; + int core_start_; + int core_end_; +}; + +// Dispatcher logic for routing WorkRequests +void DispatchWorkRequest(WorkRequest* req); + +} // namespace tdsul + +#endif // TDSUL_SRC_UTIL_THREAD_POOL_CONDUCTOR_H_ diff --git a/tpu_sync/tpudirect_storage/tests/io_queue_test.cpp b/tpu_sync/tpudirect_storage/tests/io_queue_test.cpp new file mode 100644 index 000000000..63dd1f97f --- /dev/null +++ b/tpu_sync/tpudirect_storage/tests/io_queue_test.cpp @@ -0,0 +1,244 @@ +// Copyright 2026 Google LLC. +// +// 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 "../src/def_internal.h" +#include "../src/util/IoQueue.h" +#include "tdsul/def.h" + +namespace { + +TEST(IoQueueTest, Push) { + tdsul::IoQueue queue; + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_WRITE, nullptr, + nullptr); + + // When pushed, request state is set to PENDING + queue.Push(&req1); + EXPECT_EQ(req1.state_, tdsul::RequestState::PENDING); + EXPECT_EQ(queue.Size(), 1); + + queue.Push(&req2); + EXPECT_EQ(req2.state_, tdsul::RequestState::PENDING); + EXPECT_EQ(queue.Size(), 2); + + // Verify FIFO order at the front + tdsul::WorkRequest* front_req = nullptr; + EXPECT_TRUE(queue.Front(front_req)); + EXPECT_EQ(front_req, &req1); +} + +TEST(IoQueueTest, Pop) { + tdsul::IoQueue queue; + + tdsul::WorkRequest* popped_req = nullptr; + // Pop on an empty queue should return false + EXPECT_FALSE(queue.Pop(popped_req)); + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_WRITE, nullptr, + nullptr); + + queue.Push(&req1); + queue.Push(&req2); + + // Pop requires the front request to be in FINISHED state. + // Initially req1 is PENDING, so Pop should return false. + EXPECT_FALSE(queue.Pop(popped_req)); + EXPECT_EQ(queue.Size(), 2); + + // Once req1 is marked FINISHED, Pop should succeed and remove req1. + req1.state_ = tdsul::RequestState::FINISHED; + EXPECT_TRUE(queue.Pop(popped_req)); + EXPECT_EQ(popped_req, &req1); + EXPECT_EQ(queue.Size(), 1); + + // req2 is still PENDING, so another Pop should fail. + EXPECT_FALSE(queue.Pop(popped_req)); + + // Mark req2 as FINISHED and pop it. + req2.state_ = tdsul::RequestState::FINISHED; + EXPECT_TRUE(queue.Pop(popped_req)); + EXPECT_EQ(popped_req, &req2); + EXPECT_EQ(queue.Size(), 0); + + // Now the queue is empty, Pop should fail. + EXPECT_FALSE(queue.Pop(popped_req)); +} + +TEST(IoQueueTest, Size) { + tdsul::IoQueue queue; + + EXPECT_EQ(queue.Size(), 0); + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_WRITE, nullptr, + nullptr); + tdsul::SingleWorkRequest req3(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + + queue.Push(&req1); + EXPECT_EQ(queue.Size(), 1); + + queue.Push(&req2); + EXPECT_EQ(queue.Size(), 2); + + queue.Push(&req3); + EXPECT_EQ(queue.Size(), 3); + + // Pop an element and check size decreases + req1.state_ = tdsul::RequestState::FINISHED; + tdsul::WorkRequest* popped = nullptr; + EXPECT_TRUE(queue.Pop(popped)); + EXPECT_EQ(queue.Size(), 2); + + req2.state_ = tdsul::RequestState::FINISHED; + EXPECT_TRUE(queue.Pop(popped)); + EXPECT_EQ(queue.Size(), 1); + + req3.state_ = tdsul::RequestState::FINISHED; + EXPECT_TRUE(queue.Pop(popped)); + EXPECT_EQ(queue.Size(), 0); +} + +TEST(IoQueueTest, Empty) { + tdsul::IoQueue queue; + + EXPECT_TRUE(queue.Empty()); + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + + queue.Push(&req); + EXPECT_FALSE(queue.Empty()); + + req.state_ = tdsul::RequestState::FINISHED; + tdsul::WorkRequest* popped = nullptr; + EXPECT_TRUE(queue.Pop(popped)); + EXPECT_TRUE(queue.Empty()); +} + +TEST(IoQueueTest, Front) { + tdsul::IoQueue queue; + + tdsul::WorkRequest* front_req = nullptr; + EXPECT_FALSE(queue.Front(front_req)); + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_WRITE, nullptr, + nullptr); + + queue.Push(&req1); + EXPECT_TRUE(queue.Front(front_req)); + EXPECT_EQ(front_req, &req1); + EXPECT_EQ(queue.Size(), 1); + + queue.Push(&req2); + EXPECT_TRUE(queue.Front(front_req)); + EXPECT_EQ(front_req, &req1); + EXPECT_EQ(queue.Size(), 2); +} + +TEST(IoQueueTest, TestAndSetAssigned) { + tdsul::IoQueue queue; + + EXPECT_TRUE( + queue.TestAndSetAssigned()); // Succeeded in setting, returns true + EXPECT_FALSE(queue.TestAndSetAssigned()); // Already true, returns false + + EXPECT_EQ(queue.CheckAndGetNextSchedulableRequest(), nullptr); + + // CheckAndGetNextSchedulableRequest sets is_assigned_ to false when returning + // nullptr + EXPECT_TRUE(queue.TestAndSetAssigned()); +} + +TEST(IoQueueTest, SetWorker) { + tdsul::IoQueue queue; + queue.SetWorker(nullptr); +} + +TEST(IoQueueTest, GetTotalTransferSize) { + tdsul::IoQueue queue; + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + buffer_io.handle = nullptr; + buffer_io.raw_ptr.size = 100; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + + queue.Push(&req1); + EXPECT_EQ(queue.GetTotalTransferSize(), 100); + + buffer_io.raw_ptr.size = 200; + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + + queue.Push(&req2); + EXPECT_EQ(queue.GetTotalTransferSize(), 300); + + req1.state_ = tdsul::RequestState::FINISHED; + tdsul::WorkRequest* popped = nullptr; + EXPECT_TRUE(queue.Pop(popped)); + EXPECT_EQ(queue.GetTotalTransferSize(), 200); +} + +TEST(IoQueueTest, CheckAndGetNextSchedulableRequest) { + tdsul::IoQueue queue; + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::SingleWorkRequest req1(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + tdsul::SingleWorkRequest req2(storage_io, buffer_io, TDS_OP_READ, nullptr, + nullptr); + + queue.Push(&req1); + queue.Push(&req2); + + auto req = queue.CheckAndGetNextSchedulableRequest(); + EXPECT_EQ(req, &req1); + EXPECT_EQ(req->state_, tdsul::RequestState::SCHEDULED); + + EXPECT_EQ(queue.CheckAndGetNextSchedulableRequest(), nullptr); + + // After the first one is popped, the next one can be scheduled + req1.state_ = tdsul::RequestState::FINISHED; + tdsul::WorkRequest* popped = nullptr; + EXPECT_TRUE(queue.Pop(popped)); + + auto req_next = queue.CheckAndGetNextSchedulableRequest(); + EXPECT_EQ(req_next, &req2); + EXPECT_EQ(req_next->state_, tdsul::RequestState::SCHEDULED); +} + +} // namespace diff --git a/tpu_sync/tpudirect_storage/tests/tdsul_test.cpp b/tpu_sync/tpudirect_storage/tests/tdsul_test.cpp new file mode 100644 index 000000000..0f0f886af --- /dev/null +++ b/tpu_sync/tpudirect_storage/tests/tdsul_test.cpp @@ -0,0 +1,492 @@ +// Copyright 2026 Google LLC. +// +// 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 "tdsul/tdsul.h" + +#include +#include +#include + +#include +#include +#include +#include +#include + +#include "../src/syscall_internal.h" +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "tdsul/def.h" + +namespace { + +class MockSyscall : public tdsul::Syscall { + public: + MOCK_METHOD(ssize_t, pread, (int fd, void* buf, size_t count, off_t offset), + (override)); + MOCK_METHOD(ssize_t, pwrite, + (int fd, const void* buf, size_t count, off_t offset), + (override)); + MOCK_METHOD(ssize_t, preadv, + (int fd, const struct iovec* iov, int iovcnt, off_t offset), + (override)); + MOCK_METHOD(ssize_t, pwritev, + (int fd, const struct iovec* iov, int iovcnt, off_t offset), + (override)); +}; + +class TdsulTest : public ::testing::Test { + protected: + static void SetUpTestSuite() {} +}; +// todo: add test for virtual pointer registration +TEST_F(TdsulTest, TestDramBufferRegistration) { + // Allocate host DRAM memory + const std::size_t kSize = 4096; + void* host_ptr = std::malloc(kSize); + ASSERT_NE(host_ptr, nullptr); + + tds_buffer_handle_t* handle = nullptr; + + // 1. Success case: Register DRAM buffer + tds_result_t res = + tds_buffer_register_vaddr(host_ptr, TDS_MEM_HOST, kSize, &handle); + EXPECT_EQ(res, TDS_SUCCESS); + EXPECT_NE(handle, nullptr); + + // 2. Deregister DRAM buffer + res = tds_buffer_deregister(handle); + EXPECT_EQ(res, TDS_SUCCESS); + + // 3. Error case: Register with null pointer + handle = nullptr; + res = tds_buffer_register_vaddr(nullptr, TDS_MEM_HOST, kSize, &handle); + EXPECT_EQ(res, TDS_ERROR_INVALID_PARAMETER); + EXPECT_EQ(handle, nullptr); + + std::free(host_ptr); +} + +TEST_F(TdsulTest, TestDeviceMemoryRegistrationDefault) { + tds_init(nullptr); + + const std::size_t kSize = 4096; + void* dummy_ptr = reinterpret_cast(0x12340000); + tds_buffer_handle_t* handle = nullptr; + + tds_result_t res = + tds_buffer_register_vaddr(dummy_ptr, TDS_MEM_DEVICE, kSize, &handle); + EXPECT_EQ(res, TDS_ERROR_P2P_UNSUPPORTED); + EXPECT_EQ(handle, nullptr); + tds_shutdown(); +} + +TEST_F(TdsulTest, TestDeviceMemoryRegistrationEnabled) { + tds_config_t* config = tds_config_create(); + tds_config_set_int(config, "enable_p2p", 1); + tds_init(config); + + const std::size_t kSize = 4096; + void* dummy_ptr = reinterpret_cast(0x12340000); + tds_buffer_handle_t* handle = nullptr; + + tds_result_t res = + tds_buffer_register_vaddr(dummy_ptr, TDS_MEM_DEVICE, kSize, &handle); + EXPECT_EQ(res, TDS_SUCCESS); + EXPECT_NE(handle, nullptr); + + if (handle) { + EXPECT_EQ(tds_buffer_deregister(handle), TDS_SUCCESS); + } + + // cleanup + tds_config_destroy(config); + tds_shutdown(); +} + +TEST_F(TdsulTest, TestDeviceDmaBufRegistrationDefault) { + tds_init(nullptr); + + const std::size_t kSize = 4096; + int dummy_fd = 42; + tds_buffer_handle_t* handle = nullptr; + + tds_result_t res = + tds_buffer_register_dmabuf(dummy_fd, 0, TDS_MEM_DEVICE, kSize, &handle); + EXPECT_EQ(res, TDS_ERROR_P2P_UNSUPPORTED); + EXPECT_EQ(handle, nullptr); + + tds_shutdown(); +} + +TEST_F(TdsulTest, TestDeviceDmaBufRegistrationEnabled) { + tds_config_t* config = tds_config_create(); + tds_config_set_int(config, "enable_p2p", 1); + tds_init(config); + + // Note: we can't fully mock mmap here easily for success case without custom + // mocking, but if we pass a valid file descriptor or just mock the sys call, + // wait. The dmabuf registration actually calls mmap on the fd! Let's pass + // /dev/zero so mmap succeeds. + int fd = open("/dev/zero", O_RDWR); + ASSERT_GE(fd, 0); + + const std::size_t kSize = 4096; + tds_buffer_handle_t* handle = nullptr; + + tds_result_t res = + tds_buffer_register_dmabuf(fd, 0, TDS_MEM_DEVICE, kSize, &handle); + EXPECT_EQ(res, TDS_SUCCESS); + EXPECT_NE(handle, nullptr); + + if (handle) { + EXPECT_EQ(tds_buffer_deregister(handle), TDS_SUCCESS); + } + + close(fd); + tds_config_destroy(config); + tds_shutdown(); +} + +TEST_F(TdsulTest, TestReadInterruptedBySignalEINTR) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + const std::size_t kSize = 1024; + void* buf = std::malloc(kSize); + tds_buffer_handle_t* buf_handle = nullptr; + ASSERT_EQ(tds_buffer_register_vaddr(buf, TDS_MEM_HOST, kSize, &buf_handle), + TDS_SUCCESS); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, kSize); + tds_buffer_io_t buffer_io = + tds_create_registered_buffer_io(buf_handle, 0, kSize); + + using ::testing::_; + using ::testing::Return; + + // Expect pread to be called. First time returns -1 with EINTR, second time + // returns kSize. + EXPECT_CALL(mock, pread(42, _, kSize, 0)) + .WillOnce([](int, void*, size_t, off_t) { + errno = EINTR; + return -1; + }) + .WillOnce(Return(kSize)); + + ssize_t res = tds_read(&storage_io, &buffer_io); + EXPECT_EQ(res, static_cast(kSize)); + + // Cleanup + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + EXPECT_EQ(tds_buffer_deregister(buf_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} + +TEST_F(TdsulTest, TestWriteInterruptedBySignalEINTR) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + const std::size_t kSize = 1024; + void* buf = std::malloc(kSize); + tds_buffer_handle_t* buf_handle = nullptr; + ASSERT_EQ(tds_buffer_register_vaddr(buf, TDS_MEM_HOST, kSize, &buf_handle), + TDS_SUCCESS); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, kSize); + tds_buffer_io_t buffer_io = + tds_create_registered_buffer_io(buf_handle, 0, kSize); + + using ::testing::_; + using ::testing::Return; + + // Expect pwrite to be called. First time returns -1 with EINTR, second time + // returns kSize. + EXPECT_CALL(mock, pwrite(42, _, kSize, 0)) + .WillOnce([](int, const void*, size_t, off_t) { + errno = EINTR; + return -1; + }) + .WillOnce(Return(kSize)); + + ssize_t res = tds_write(&storage_io, &buffer_io); + EXPECT_EQ(res, static_cast(kSize)); + + // Cleanup + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + EXPECT_EQ(tds_buffer_deregister(buf_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} + +TEST_F(TdsulTest, ReadvEINTRRetry) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + const std::size_t kSize = 1024; + void* buf = std::malloc(kSize); + tds_buffer_io_t buffer_io = tds_create_raw_buffer_io(buf, kSize); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, kSize); + + using ::testing::_; + using ::testing::Return; + + EXPECT_CALL(mock, preadv(42, _, 1, 0)) + .WillOnce([](int, const struct iovec*, int, off_t) { + errno = EINTR; + return -1; + }) + .WillOnce(Return(kSize)); + + ssize_t res = tds_readv(&storage_io, &buffer_io, 1); + EXPECT_EQ(res, static_cast(kSize)); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} + +TEST_F(TdsulTest, WritevEINTRRetry) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + const std::size_t kSize = 1024; + void* buf = std::malloc(kSize); + tds_buffer_io_t buffer_io = tds_create_raw_buffer_io(buf, kSize); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, kSize); + + using ::testing::_; + using ::testing::Return; + + EXPECT_CALL(mock, pwritev(42, _, 1, 0)) + .WillOnce([](int, const struct iovec*, int, off_t) { + errno = EINTR; + return -1; + }) + .WillOnce(Return(kSize)); + + ssize_t res = tds_writev(&storage_io, &buffer_io, 1); + EXPECT_EQ(res, static_cast(kSize)); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} + +TEST_F(TdsulTest, VectoredIoSizeOverflow) { + void* buf = std::malloc(1024); + + tds_buffer_io_t buffer_iov[2]; + buffer_iov[0] = tds_create_raw_buffer_io(buf, SSIZE_MAX - 100); + buffer_iov[1] = tds_create_raw_buffer_io(buf, 200); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, 1024); + + errno = 0; + ssize_t res = tds_readv(&storage_io, buffer_iov, 2); + EXPECT_EQ(res, -1); + EXPECT_EQ(errno, EINVAL); + + errno = 0; + res = tds_writev(&storage_io, buffer_iov, 2); + EXPECT_EQ(res, -1); + EXPECT_EQ(errno, EINVAL); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); +} + +TEST_F(TdsulTest, RegisteredBufferBoundsCheck) { + void* buf = std::malloc(1024); + + tds_buffer_handle_t* handle = nullptr; + ASSERT_EQ(tds_buffer_register_vaddr(buf, TDS_MEM_HOST, 1024, &handle), + TDS_SUCCESS); + + // offset + size > 1024 + tds_buffer_io_t buffer_io = tds_create_registered_buffer_io(handle, 500, 600); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, 1024); + + errno = 0; + ssize_t res = tds_read(&storage_io, &buffer_io); + EXPECT_EQ(res, -1); + EXPECT_EQ(errno, EINVAL); + + EXPECT_EQ(tds_buffer_deregister(handle), TDS_SUCCESS); + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); +} + +TEST_F(TdsulTest, VectoredIoInvalidArgs) { + void* buf = std::malloc(1024); + tds_buffer_io_t bio = tds_create_raw_buffer_io(buf, 1024); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, 1024); + + // Null storage_io + EXPECT_EQ(tds_readv(nullptr, &bio, 1), -1); + // Null storage_io->handle + tds_storage_io_t null_handle_io = storage_io; + null_handle_io.handle = nullptr; + EXPECT_EQ(tds_readv(&null_handle_io, &bio, 1), -1); + // Null buffer_iov + EXPECT_EQ(tds_readv(&storage_io, nullptr, 1), -1); + // iovcnt <= 0 + EXPECT_EQ(tds_readv(&storage_io, &bio, 0), -1); + // iovcnt > IOV_MAX + std::vector huge_iov(IOV_MAX + 1, bio); + EXPECT_EQ(tds_readv(&storage_io, huge_iov.data(), IOV_MAX + 1), -1); + + // Same for writev + EXPECT_EQ(tds_writev(nullptr, &bio, 1), -1); + EXPECT_EQ(tds_writev(&null_handle_io, &bio, 1), -1); + EXPECT_EQ(tds_writev(&storage_io, nullptr, 1), -1); + EXPECT_EQ(tds_writev(&storage_io, &bio, -1), -1); + EXPECT_EQ(tds_writev(&storage_io, huge_iov.data(), IOV_MAX + 1), -1); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); +} + +TEST_F(TdsulTest, VectoredIoZeroSize) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + void* buf = std::malloc(1024); + tds_buffer_io_t bio = tds_create_raw_buffer_io(buf, 0); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, 1024); + + // Expect 0 without calling preadv/pwritev + EXPECT_CALL(mock, + preadv(::testing::_, ::testing::_, ::testing::_, ::testing::_)) + .Times(0); + EXPECT_CALL(mock, + pwritev(::testing::_, ::testing::_, ::testing::_, ::testing::_)) + .Times(0); + + EXPECT_EQ(tds_readv(&storage_io, &bio, 1), 0); + EXPECT_EQ(tds_writev(&storage_io, &bio, 1), 0); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} + +TEST_F(TdsulTest, VectoredIoSuccess) { + MockSyscall mock; + tdsul::Syscall* original_syscall = tdsul::g_syscall; + tdsul::g_syscall = &mock; + + const std::size_t kSize = 1024; + void* buf = std::malloc(kSize); + tds_buffer_io_t buffer_io[2]; + buffer_io[0] = tds_create_raw_buffer_io(buf, kSize / 2); + buffer_io[1] = + tds_create_raw_buffer_io(static_cast(buf) + kSize / 2, kSize / 2); + + tds_storage_handle_t* storage_handle = nullptr; + tds_storage_descr_t storage_descr; + storage_descr.uri = "fd://42"; + storage_descr.options = nullptr; + ASSERT_EQ(tds_storage_handle_register(&storage_descr, &storage_handle), + TDS_SUCCESS); + + tds_storage_io_t storage_io = tds_create_file_io(storage_handle, 0, kSize); + + using ::testing::_; + using ::testing::Return; + + EXPECT_CALL(mock, preadv(42, _, 2, 0)).WillOnce(Return(kSize)); + ssize_t res = tds_readv(&storage_io, buffer_io, 2); + EXPECT_EQ(res, static_cast(kSize)); + + EXPECT_CALL(mock, pwritev(42, _, 2, 0)).WillOnce(Return(kSize)); + res = tds_writev(&storage_io, buffer_io, 2); + EXPECT_EQ(res, static_cast(kSize)); + + EXPECT_EQ(tds_storage_handle_deregister(storage_handle), TDS_SUCCESS); + std::free(buf); + tdsul::g_syscall = original_syscall; +} +// TODO: add tests for tdsul async IO after the logic is finished. We shall +// specically check: +// (1) whether the FIFO order is guaranteed, w/ or w/o batched IO. +// (2) whether concurrency submission is properly handled. +} // namespace diff --git a/tpu_sync/tpudirect_storage/tests/thread_pool_conductor_test.cpp b/tpu_sync/tpudirect_storage/tests/thread_pool_conductor_test.cpp new file mode 100644 index 000000000..16b3c9ff1 --- /dev/null +++ b/tpu_sync/tpudirect_storage/tests/thread_pool_conductor_test.cpp @@ -0,0 +1,125 @@ +// Copyright 2026 Google LLC. +// +// 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 "../src/def_internal.h" +#include "../src/syscall_internal.h" +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "tdsul/def.h" +#include "util/IoConductor.h" +#include "util/ThreadPoolConductor.h" + +namespace tdsul { +extern std::unique_ptr g_io_conductor; +} + +namespace { + +class MockSyscall : public tdsul::Syscall { + public: + MOCK_METHOD(ssize_t, pread, (int fd, void* buf, size_t count, off_t offset), + (override)); + MOCK_METHOD(ssize_t, pwrite, + (int fd, const void* buf, size_t count, off_t offset), + (override)); +}; + +class ThreadPoolConductorTest : public ::testing::Test { + protected: + static void SetUpTestSuite() {} + + void SetUp() override { + original_syscall_ = tdsul::g_syscall; + tdsul::g_syscall = &mock_syscall_; + } + + void TearDown() override { tdsul::g_syscall = original_syscall_; } + + MockSyscall mock_syscall_; + tdsul::Syscall* original_syscall_; +}; + +TEST_F(ThreadPoolConductorTest, TestWorkerThreadDispatch) { + // Set up mock expectations. Since the worker thread doesn't yet call + // pread/pwrite, we just allow it to be called any number of times. + // When the logic is added to ThreadPoolConductor::Run, these tests can be + // tightened. + using ::testing::_; + EXPECT_CALL(mock_syscall_, pread(_, _, _, _)).Times(testing::AnyNumber()); + EXPECT_CALL(mock_syscall_, pwrite(_, _, _, _)).Times(testing::AnyNumber()); + + tdsul::ThreadPoolConductor pool(2); // 2 threads + pool.Start(); + + // Each request carries a future so the test can prove the worker threads + // actually executed the work. Asserting only "no crash" previously allowed a + // pool whose threads failed to start to pass while silently dropping (and + // leaking) every dispatched request. + const int kNumTasks = 10; + std::vector futures(kNumTasks); + for (auto& future : futures) { + future.status.status = TDS_PENDING; + } + + for (int i = 0; i < kNumTasks; ++i) { + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + tdsul::WorkRequest* req = new tdsul::SingleWorkRequest( + storage_io, buffer_io, TDS_OP_READ, nullptr, &futures[i]); + + pool.Dispatch(req); + } + + // Stop() waits for all tasks in the queues to be processed (and deleted). + pool.Stop(); + + // Every request must have reached a worker and completed. The heap checker + // independently verifies that each `SingleWorkRequest` was freed. + for (int i = 0; i < kNumTasks; ++i) { + EXPECT_EQ(futures[i].status.status, TDS_SUCCESS) + << "Request " << i << " was never executed by a worker thread"; + } +} + +// A pool whose worker threads were never started must not silently retain (and +// leak) dispatched work. `Stop()` drains any queued requests and marks their +// futures failed, so callers observe an error instead of waiting forever. +TEST_F(ThreadPoolConductorTest, StopReleasesRequestsQueuedBeforeStart) { + tdsul::ThreadPoolConductor pool(2); + // Deliberately no Start(): the worker threads are constructed but idle, + // mirroring the state left behind when pthread_create fails outright. + + tds_io_future future{}; + future.status.status = TDS_PENDING; + + tds_storage_io_t storage_io{}; + tds_buffer_io_t buffer_io{}; + pool.Dispatch(new tdsul::SingleWorkRequest(storage_io, buffer_io, TDS_OP_READ, + nullptr, &future)); + + pool.Stop(); + + EXPECT_NE(future.status.status, TDS_PENDING) + << "Abandoned request left its future pending forever"; +} + +} // namespace