diff --git a/src/paimon/common/utils/arrow/arrow_input_stream_adapter.cpp b/src/paimon/common/utils/arrow/arrow_input_stream_adapter.cpp index 0d916d480..50510a4b2 100644 --- a/src/paimon/common/utils/arrow/arrow_input_stream_adapter.cpp +++ b/src/paimon/common/utils/arrow/arrow_input_stream_adapter.cpp @@ -40,6 +40,19 @@ arrow::Status ValidateArrowIoRange(int64_t value, const char* name) { return arrow::Status::OK(); } +struct BufferWithMemoryPool { + std::shared_ptr pool; + std::shared_ptr buffer; +}; + +std::shared_ptr KeepMemoryPoolAlive(std::shared_ptr buffer, + const std::shared_ptr& pool) { + auto holder = + std::make_shared(BufferWithMemoryPool{pool, std::move(buffer)}); + auto* buffer_ptr = holder->buffer.get(); + return std::shared_ptr(std::move(holder), buffer_ptr); +} + } // namespace ArrowInputStreamAdapter::ArrowInputStreamAdapter( @@ -79,7 +92,7 @@ arrow::Result> ArrowInputStreamAdapter::Read(int6 if (read_bytes < nbytes) { ARROW_RETURN_NOT_OK(buffer->Resize(read_bytes)); } - return std::shared_ptr(std::move(buffer)); + return KeepMemoryPoolAlive(std::shared_ptr(std::move(buffer)), pool_); } arrow::Result ArrowInputStreamAdapter::ReadAt(int64_t position, int64_t nbytes, @@ -104,7 +117,7 @@ arrow::Result> ArrowInputStreamAdapter::ReadAt(in if (read_bytes < nbytes) { ARROW_RETURN_NOT_OK(buffer->Resize(read_bytes)); } - return std::shared_ptr(std::move(buffer)); + return KeepMemoryPoolAlive(std::shared_ptr(std::move(buffer)), pool_); } arrow::Future> ArrowInputStreamAdapter::ReadAsync( @@ -128,6 +141,7 @@ arrow::Future> ArrowInputStreamAdapter::ReadAsync return fut; } std::shared_ptr buffer = std::move(buffer_result).ValueUnsafe(); + buffer = KeepMemoryPoolAlive(std::move(buffer), pool_); std::shared_ptr> storage_read_bytes = storage_read_bytes_; input_stream_->ReadAsync( reinterpret_cast(buffer->mutable_data()), nbytes, position, diff --git a/src/paimon/common/utils/arrow/arrow_stream_adapter_test.cpp b/src/paimon/common/utils/arrow/arrow_stream_adapter_test.cpp index c0b14395f..12f584904 100644 --- a/src/paimon/common/utils/arrow/arrow_stream_adapter_test.cpp +++ b/src/paimon/common/utils/arrow/arrow_stream_adapter_test.cpp @@ -15,8 +15,11 @@ */ #include +#include +#include #include #include +#include #include "arrow/api.h" #include "arrow/io/type_fwd.h" @@ -32,6 +35,79 @@ namespace paimon::test { +namespace { + +constexpr char kTestPayload[] = "data"; +constexpr int64_t kTestSize = sizeof(kTestPayload) - 1; + +class DeferredInputStream : public InputStream { + public: + Status Seek(int64_t, SeekOrigin) override { + return Status::OK(); + } + + Result GetPos() const override { + return 0; + } + + Result Read(char* buffer, int64_t size) override { + if (size != kTestSize) { + return Status::Invalid("unexpected read size"); + } + std::memcpy(buffer, kTestPayload, kTestSize); + return kTestSize; + } + + Result Read(char* buffer, int64_t size, int64_t) override { + return Read(buffer, size); + } + + void ReadAsync(char* buffer, int64_t size, int64_t, + std::function&& callback) override { + buffer_ = buffer; + size_ = size; + callback_ = std::move(callback); + } + + Status Complete() { + if (!callback_) { + return Status::Invalid("async request was not started"); + } + if (size_ != kTestSize) { + return Status::Invalid("unexpected async read size"); + } + std::memcpy(buffer_, kTestPayload, kTestSize); + auto callback = std::move(callback_); + callback(Status::OK()); + return Status::OK(); + } + + Status Close() override { + return Status::OK(); + } + + Result GetUri() const override { + return std::string("test://input"); + } + + Result Length() const override { + return kTestSize; + } + + private: + char* buffer_ = nullptr; + int64_t size_ = 0; + std::function callback_; +}; + +std::shared_ptr CreateAdapter( + const std::shared_ptr& stream, + const std::shared_ptr& pool) { + return std::make_shared(stream, kTestSize, pool); +} + +} // namespace + TEST(ArrowStreamAdapterTest, TestInputAndOutputStream) { auto test_root_dir = UniqueTestDirectory::Create(); ASSERT_TRUE(test_root_dir); @@ -88,4 +164,48 @@ TEST(ArrowStreamAdapterTest, TestInputAndOutputStream) { ASSERT_TRUE(in_stream->closed()); } +TEST(ArrowStreamAdapterTest, TestReadKeepsMemoryPoolAliveUntilBufferReleased) { + auto stream = std::make_shared(); + std::shared_ptr pool(GetArrowPool(GetDefaultPool())); + auto adapter = CreateAdapter(stream, pool); + ASSERT_NE(adapter, nullptr); + + std::shared_ptr buffer = adapter->Read(kTestSize).ValueOrDie(); + adapter.reset(); + pool.reset(); + + ASSERT_EQ(buffer->ToString(), kTestPayload); + buffer.reset(); +} + +TEST(ArrowStreamAdapterTest, TestReadAtKeepsMemoryPoolAliveUntilBufferReleased) { + auto stream = std::make_shared(); + std::shared_ptr pool(GetArrowPool(GetDefaultPool())); + auto adapter = CreateAdapter(stream, pool); + ASSERT_NE(adapter, nullptr); + + std::shared_ptr buffer = adapter->ReadAt(0, kTestSize).ValueOrDie(); + adapter.reset(); + pool.reset(); + + ASSERT_EQ(buffer->ToString(), kTestPayload); + buffer.reset(); +} + +TEST(ArrowStreamAdapterTest, TestAsyncReadKeepsMemoryPoolAliveUntilBufferReleased) { + auto stream = std::make_shared(); + std::shared_ptr pool(GetArrowPool(GetDefaultPool())); + auto adapter = CreateAdapter(stream, pool); + ASSERT_NE(adapter, nullptr); + + auto future = adapter->ReadAsync(arrow::io::default_io_context(), 0, kTestSize); + adapter.reset(); + pool.reset(); + + ASSERT_OK(stream->Complete()); + std::shared_ptr buffer = future.MoveResult().ValueOrDie(); + ASSERT_EQ(buffer->ToString(), kTestPayload); + buffer.reset(); +} + } // namespace paimon::test