This is an automated email from the ASF dual-hosted git repository.
marin-ma pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/gluten.git
The following commit(s) were added to refs/heads/main by this push:
new cf7650496d [GLUTEN-12985][VL] Fix multi-window page deserialization in
rss-sort shuffle reader (#12984)
cf7650496d is described below
commit cf7650496df73f79e760de84e20bdea1e98d26ab
Author: Kuo Zhao <[email protected]>
AuthorDate: Mon Sep 21 16:47:43 2026 +0800
[GLUTEN-12985][VL] Fix multi-window page deserialization in rss-sort
shuffle reader (#12984)
---
cpp/velox/shuffle/GlutenByteStream.h | 271 ------------------------------
cpp/velox/shuffle/VeloxShuffleReader.cc | 261 +++++++++++++++++++++++++---
cpp/velox/shuffle/VeloxShuffleReader.h | 6 +-
cpp/velox/tests/VeloxShuffleReaderTest.cc | 164 +++++++++++++++++-
4 files changed, 402 insertions(+), 300 deletions(-)
diff --git a/cpp/velox/shuffle/GlutenByteStream.h
b/cpp/velox/shuffle/GlutenByteStream.h
deleted file mode 100644
index 8085e8a743..0000000000
--- a/cpp/velox/shuffle/GlutenByteStream.h
+++ /dev/null
@@ -1,271 +0,0 @@
-/*
- * Licensed to the Apache Software Foundation (ASF) under one or more
- * contributor license agreements. See the NOTICE file distributed with
- * this work for additional information regarding copyright ownership.
- * The ASF licenses this file to You 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.
- */
-
-// TODO: wait to delete after rss sort reader refactored.
-#include "velox/common/memory/ByteStream.h"
-
-namespace facebook::velox {
-
-class GlutenByteInputStream : public ByteInputStream {
- protected:
- /// TODO Remove after refactoring SpillInput.
- GlutenByteInputStream() {}
-
- public:
- explicit GlutenByteInputStream(std::vector<ByteRange> ranges) {
- ranges_ = std::move(ranges);
- VELOX_CHECK(!ranges_.empty());
- current_ = &ranges_[0];
- }
-
- /// Disable copy constructor.
- GlutenByteInputStream(const GlutenByteInputStream&) = delete;
-
- /// Disable copy assignment operator.
- GlutenByteInputStream& operator=(const GlutenByteInputStream& other) =
delete;
-
- /// Enable move constructor.
- GlutenByteInputStream(GlutenByteInputStream&& other) noexcept = delete;
-
- /// Enable move assignment operator.
- GlutenByteInputStream& operator=(GlutenByteInputStream&& other) noexcept {
- if (this != &other) {
- ranges_ = std::move(other.ranges_);
- current_ = other.current_;
- other.current_ = nullptr;
- }
- return *this;
- }
-
- /// TODO Remove after refactoring SpillInput.
- virtual ~GlutenByteInputStream() = default;
-
- std::vector<ByteRange> ranges_;
-
- /// Returns total number of bytes available in the stream.
- size_t size() const {
- size_t total = 0;
- for (const auto& range : ranges_) {
- total += range.size;
- }
- return total;
- }
-
- /// Returns true if all input has been read.
- ///
- /// TODO: Remove 'virtual' after refactoring SpillInput.
- virtual bool atEnd() const {
- if (!current_) {
- return false;
- }
- if (current_->position < current_->size) {
- return false;
- }
-
- VELOX_CHECK(current_ >= ranges_.data() && current_ <= &ranges_.back());
- return current_ == &ranges_.back();
- }
-
- /// Returns current position (number of bytes from the start) in the stream.
- std::streampos tellp() const {
- if (ranges_.empty()) {
- return 0;
- }
- VELOX_DCHECK_NOT_NULL(current_);
- int64_t size = 0;
- for (auto& range : ranges_) {
- if (&range == current_) {
- return current_->position + size;
- }
- size += range.size;
- }
- VELOX_FAIL("GlutenByteInputStream 'current_' is not in 'ranges_'.");
- }
-
- /// Moves current position to specified one.
- void seekp(std::streampos position) {
- if (ranges_.empty() && position == 0) {
- return;
- }
- int64_t toSkip = position;
- for (auto& range : ranges_) {
- if (toSkip <= range.size) {
- current_ = ⦥
- current_->position = toSkip;
- return;
- }
- toSkip -= range.size;
- }
- static_assert(sizeof(std::streamsize) <= sizeof(long long));
- VELOX_FAIL("Seeking past end of GlutenByteInputStream: {}",
static_cast<long long>(position));
- }
-
- /// Returns the remaining size left from current reading position.
- size_t remainingSize() const {
- if (ranges_.empty()) {
- return 0;
- }
- const auto* lastRange = &ranges_[ranges_.size() - 1];
- auto cur = current_;
- size_t total = cur->size - cur->position;
- while (++cur <= lastRange) {
- total += cur->size;
- }
- return total;
- }
-
- std::string toString() const {
- std::stringstream oss;
- oss << ranges_.size() << " ranges (position/size) [";
- for (const auto& range : ranges_) {
- oss << "(" << range.position << "/" << range.size << (&range == current_
? " current" : "") << ")";
- if (&range != &ranges_.back()) {
- oss << ",";
- }
- }
- oss << "]";
- return oss.str();
- }
-
- uint8_t readByte() {
- if (current_->position < current_->size) {
- return current_->buffer[current_->position++];
- }
- next();
- return readByte();
- }
-
- void readBytes(uint8_t* bytes, int32_t size) {
- VELOX_CHECK_GE(size, 0, "Attempting to read negative number of bytes");
- int32_t offset = 0;
- for (;;) {
- int32_t available = current_->size - current_->position;
- int32_t numUsed = std::min(available, size);
- simd::memcpy(bytes + offset, current_->buffer + current_->position,
numUsed);
- offset += numUsed;
- size -= numUsed;
- current_->position += numUsed;
- if (!size) {
- return;
- }
- next();
- }
- }
-
- template <typename T>
- T read() {
- if (current_->position + sizeof(T) <= current_->size) {
- current_->position += sizeof(T);
- return *reinterpret_cast<const T*>(current_->buffer + current_->position
- sizeof(T));
- }
- // The number straddles two buffers. We read byte by byte and make
- // a little-endian uint64_t. The bytes can be cast to any integer
- // or floating point type since the wire format has the machine byte order.
- static_assert(sizeof(T) <= sizeof(uint64_t));
- uint64_t value = 0;
- for (int32_t i = 0; i < sizeof(T); ++i) {
- value |= static_cast<uint64_t>(readByte()) << (i * 8);
- }
- return *reinterpret_cast<const T*>(&value);
- }
-
- template <typename Char>
- void readBytes(Char* data, int32_t size) {
- readBytes(reinterpret_cast<uint8_t*>(data), size);
- }
-
- /// Returns a view over the read buffer for up to 'size' next
- /// bytes. The size of the value may be less if the current byte
- /// range ends within 'size' bytes from the current position. The
- /// size will be 0 if at end.
- std::string_view nextView(int64_t size) {
- VELOX_CHECK_GE(size, 0, "Attempting to view negative number of bytes");
- if (current_->position == current_->size) {
- if (current_ == &ranges_.back()) {
- return std::string_view(nullptr, 0);
- }
- next();
- }
- VELOX_CHECK(current_->size);
- auto position = current_->position;
- auto viewSize = std::min(current_->size - current_->position, size);
- current_->position += viewSize;
- return std::string_view(reinterpret_cast<char*>(current_->buffer) +
position, viewSize);
- }
-
- void skip(int32_t size) {
- VELOX_CHECK_GE(size, 0, "Attempting to skip negative number of bytes");
- for (;;) {
- int32_t available = current_->size - current_->position;
- int32_t numUsed = std::min(available, size);
- size -= numUsed;
- current_->position += numUsed;
- if (!size) {
- return;
- }
- next();
- }
- }
-
- protected:
- /// Sets 'current_' to point to the next range of input. // The
- /// input is consecutive ByteRanges in 'ranges_' for the base class
- /// but any view over external buffers can be made by specialization.
- ///
- /// TODO: Remove 'virtual' after refactoring SpillInput.
- virtual void next(bool throwIfPastEnd = true) {
- VELOX_CHECK(current_ >= &ranges_[0]);
- size_t position = current_ - &ranges_[0];
- VELOX_CHECK_LT(position, ranges_.size());
- if (position == ranges_.size() - 1) {
- if (throwIfPastEnd) {
- VELOX_FAIL("Reading past end of GlutenByteInputStream");
- }
- return;
- }
- ++current_;
- current_->position = 0;
- }
-
- // TODO: Remove after refactoring SpillInput.
- const std::vector<ByteRange>& ranges() const {
- return ranges_;
- }
-
- // TODO: Remove after refactoring SpillInput.
- void setRange(ByteRange range) {
- ranges_.resize(1);
- ranges_[0] = range;
- current_ = ranges_.data();
- }
-};
-
-template <>
-inline Timestamp GlutenByteInputStream::read<Timestamp>() {
- Timestamp value;
- readBytes(reinterpret_cast<uint8_t*>(&value), sizeof(value));
- return value;
-}
-
-template <>
-inline int128_t GlutenByteInputStream::read<int128_t>() {
- int128_t value;
- readBytes(reinterpret_cast<uint8_t*>(&value), sizeof(value));
- return value;
-}
-
-} // namespace facebook::velox
diff --git a/cpp/velox/shuffle/VeloxShuffleReader.cc
b/cpp/velox/shuffle/VeloxShuffleReader.cc
index 6c2e101b40..478d619773 100644
--- a/cpp/velox/shuffle/VeloxShuffleReader.cc
+++ b/cpp/velox/shuffle/VeloxShuffleReader.cc
@@ -22,20 +22,24 @@
#include "compute/VeloxBackend.h"
#include "memory/VeloxColumnarBatch.h"
-#include "shuffle/GlutenByteStream.h"
#include "shuffle/Payload.h"
#include "shuffle/Utils.h"
#include "utils/Common.h"
#include "utils/Timer.h"
#include "utils/VeloxArrowUtils.h"
+#include "velox/common/memory/ByteStream.h"
#include "velox/row/CompactRow.h"
+#include "velox/serializers/PrestoHeader.h"
#include "velox/serializers/PrestoSerializer.h"
+#include "velox/serializers/PrestoSerializerSerializationUtils.h"
#include "velox/vector/ComplexVector.h"
#include "velox/vector/FlatVector.h"
#include "velox/vector/arrow/Bridge.h"
#include <algorithm>
+#include <array>
+#include <sstream>
#include "VeloxGpuAsyncShuffleReader.h"
#include "config/VeloxConfig.h"
@@ -826,49 +830,111 @@ void VeloxSortShuffleReaderDeserializer::readNextRow() {
++cachedRows_;
}
-class VeloxRssSortShuffleReaderDeserializer::VeloxInputStream : public
facebook::velox::GlutenByteInputStream {
+// A single-window refill stream: each next() overwrites the sole range with up
+// to buffer_ capacity bytes from the underlying InputStream. Because earlier
+// windows are physically discarded on refill, callers can always read forward
+// and rewind within the current window, but cannot revisit bytes from prior
+// windows — seekp() fails fast instead of reading overwritten data.
+class VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream :
public facebook::velox::ByteInputStream {
public:
- VeloxInputStream(std::shared_ptr<arrow::io::InputStream> input,
facebook::velox::BufferPtr buffer);
+ RssSortShuffleReaderInputStream(std::shared_ptr<arrow::io::InputStream>
input, facebook::velox::BufferPtr buffer);
bool hasNext();
- void next(bool throwIfPastEnd) override;
+ /// Refills the window from the underlying InputStream. Throws if
+ /// 'throwIfPastEnd' and the stream is already exhausted.
+ void next(bool throwIfPastEnd = true);
size_t remainingSize() const override;
+ size_t size() const override {
+ return atEnd_ ? totalBytesRead_ : std::numeric_limits<size_t>::max();
+ }
+
+ bool atEnd() const override {
+ return atEnd_;
+ }
+
+ std::streampos tellp() const override;
+
+ void seekp(std::streampos position) override;
+
+ uint8_t readByte() override;
+
+ void readBytes(uint8_t* bytes, int32_t size) override;
+
+ std::string_view nextView(int64_t size) override;
+
+ void skip(int32_t size) override;
+
+ std::string toString() const override;
+
+ int32_t remainingInWindow() const {
+ if (ranges_.empty()) {
+ return 0;
+ }
+ return ranges_[0].size - ranges_[0].position;
+ }
+
+ uint8_t* data() const {
+ return ranges_.empty() ? nullptr : ranges_[0].buffer + ranges_[0].position;
+ }
+
+ void advance(int32_t n) {
+ VELOX_CHECK(!ranges_.empty() && ranges_[0].position + n <=
ranges_[0].size);
+ ranges_[0].position += n;
+ }
+
+ private:
+ void setRange(ByteRange range) {
+ ranges_.resize(1);
+ ranges_[0] = range;
+ current_ = ranges_.data();
+ }
+
std::shared_ptr<arrow::io::InputStream> in_;
const facebook::velox::BufferPtr buffer_;
uint64_t offset_ = -1;
+ uint64_t totalBytesRead_ = 0;
+ bool atEnd_ = false;
+ std::vector<ByteRange> ranges_;
};
-VeloxRssSortShuffleReaderDeserializer::VeloxInputStream::VeloxInputStream(
+VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::RssSortShuffleReaderInputStream(
std::shared_ptr<arrow::io::InputStream> input,
facebook::velox::BufferPtr buffer)
: in_(std::move(input)), buffer_(std::move(buffer)) {
next(false);
}
-bool VeloxRssSortShuffleReaderDeserializer::VeloxInputStream::hasNext() {
+bool
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::hasNext()
{
if (offset_ == 0) {
return false;
}
- if (ranges()[0].position >= ranges()[0].size) {
+ if (ranges_[0].position >= ranges_[0].size) {
next(false);
return offset_ != 0;
}
return true;
}
-void VeloxRssSortShuffleReaderDeserializer::VeloxInputStream::next(bool
throwIfPastEnd) {
+void
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::next(bool
throwIfPastEnd) {
const uint32_t readBytes = buffer_->capacity();
offset_ = 0;
GLUTEN_ASSIGN_OR_THROW(int64_t realBytes, in_->Read(readBytes,
buffer_->asMutable<char>()));
if (realBytes > 0) {
offset_ = realBytes;
+ totalBytesRead_ += realBytes;
+ atEnd_ = false;
setRange({buffer_->asMutable<uint8_t>(), static_cast<int32_t>(realBytes),
0});
- } else if (throwIfPastEnd) {
- VELOX_FAIL(
- "Reading past end of
VeloxRssSortShuffleReaderDeserializer::VeloxInputStream, real bytes = {}",
realBytes);
+ } else {
+ atEnd_ = true;
+ if (throwIfPastEnd) {
+ VELOX_FAIL(
+ "Reading past end of RssSortShuffleReaderInputStream, real bytes =
{}, totalBytesRead = {}",
+ realBytes,
+ totalBytesRead_);
+ }
}
}
@@ -909,19 +975,14 @@ std::shared_ptr<ColumnarBatch>
VeloxRssSortShuffleReaderDeserializer::next() {
ScopedTimer timer(&deserializeTime_);
- RowVectorPtr rowVector;
- VectorStreamGroup::read(
- in_.get(), memoryManager_->getLeafMemoryPool().get(), rowType_, serde_,
&rowVector, &serdeOptions_);
+ auto rowVector = readPage();
if (rowVector->size() >= batchSize_) {
return std::make_shared<VeloxColumnarBatch>(std::move(rowVector));
}
while (rowVector->size() < batchSize_ && in_->hasNext()) {
- RowVectorPtr rowVectorTemp;
- VectorStreamGroup::read(
- in_.get(), memoryManager_->getLeafMemoryPool().get(), rowType_,
serde_, &rowVectorTemp, &serdeOptions_);
- rowVector->append(rowVectorTemp.get());
+ rowVector->append(readPage().get());
}
return std::make_shared<VeloxColumnarBatch>(std::move(rowVector));
@@ -948,13 +1009,173 @@ void
VeloxRssSortShuffleReaderDeserializer::loadNextStream() {
constexpr uint64_t kMaxReadBufferSize = (1 << 20) -
AlignedBuffer::kPaddedSize;
auto buffer = AlignedBuffer::allocate<char>(kMaxReadBufferSize,
memoryManager_->getLeafMemoryPool().get());
- in_ = std::make_unique<VeloxInputStream>(std::move(arrowIn_),
std::move(buffer));
+ in_ = std::make_unique<RssSortShuffleReaderInputStream>(std::move(arrowIn_),
std::move(buffer));
}
-size_t
VeloxRssSortShuffleReaderDeserializer::VeloxInputStream::remainingSize() const {
+size_t
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::remainingSize()
const {
return std::numeric_limits<unsigned long>::max();
}
+uint8_t
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::readByte()
{
+ if (current_->position < current_->size) {
+ return current_->buffer[current_->position++];
+ }
+ next();
+ return readByte();
+}
+
+void
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::readBytes(uint8_t*
bytes, int32_t size) {
+ VELOX_CHECK_GE(size, 0, "Attempting to read negative number of bytes");
+ int32_t offset = 0;
+ for (;;) {
+ int32_t available = current_->size - current_->position;
+ int32_t numUsed = std::min(available, size);
+ simd::memcpy(bytes + offset, current_->buffer + current_->position,
numUsed);
+ offset += numUsed;
+ size -= numUsed;
+ current_->position += numUsed;
+ if (!size) {
+ return;
+ }
+ next();
+ }
+}
+
+void
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::skip(int32_t
size) {
+ VELOX_CHECK_GE(size, 0, "Attempting to skip negative number of bytes");
+ for (;;) {
+ int32_t available = current_->size - current_->position;
+ int32_t numUsed = std::min(available, size);
+ size -= numUsed;
+ current_->position += numUsed;
+ if (!size) {
+ return;
+ }
+ next();
+ }
+}
+
+std::string
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::toString()
const {
+ std::stringstream oss;
+ oss << ranges_.size() << " ranges (position/size) [";
+ for (const auto& range : ranges_) {
+ oss << "(" << range.position << "/" << range.size << (&range == current_ ?
" current" : "") << ")";
+ if (&range != &ranges_.back()) {
+ oss << ",";
+ }
+ }
+ oss << "]";
+ return oss.str();
+}
+
+std::string_view
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::nextView(int64_t
size) {
+ VELOX_CHECK_GE(size, 0, "Attempting to view negative number of bytes");
+ if (ranges_.empty()) {
+ return std::string_view(nullptr, 0);
+ }
+ if (ranges_[0].position == ranges_[0].size) {
+ // Current window is exhausted. For single-window refill streams, next()
+ // overwrites the window with fresh data, so attempt refill before
+ // reporting end-of-stream.
+ next(false);
+ if (ranges_.empty() || ranges_[0].position == ranges_[0].size) {
+ return std::string_view(nullptr, 0);
+ }
+ }
+ VELOX_DCHECK(ranges_[0].size > 0);
+ const int32_t position = ranges_[0].position;
+ const int64_t viewSize = std::min<int64_t>(ranges_[0].size -
ranges_[0].position, size);
+ ranges_[0].position += static_cast<int32_t>(viewSize);
+ return std::string_view(reinterpret_cast<char*>(ranges_[0].buffer) +
position, viewSize);
+}
+
+std::streampos
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::tellp()
const {
+ if (ranges_.empty()) {
+ return 0;
+ }
+ return static_cast<std::streampos>(static_cast<int64_t>(totalBytesRead_) -
(ranges_[0].size - ranges_[0].position));
+}
+
+void
VeloxRssSortShuffleReaderDeserializer::RssSortShuffleReaderInputStream::seekp(std::streampos
position) {
+ if (ranges_.empty() && position == 0) {
+ return;
+ }
+ VELOX_CHECK(!ranges_.empty(), "Cannot seek an empty
RssSortShuffleReaderInputStream");
+ const int64_t windowStart = static_cast<int64_t>(totalBytesRead_) -
ranges_[0].size;
+ const int64_t windowEnd = static_cast<int64_t>(totalBytesRead_);
+ const int64_t target = static_cast<int64_t>(position);
+ VELOX_CHECK(
+ target >= windowStart && target <= windowEnd,
+ "RssSortShuffleReaderInputStream::seekp({}) is outside the resident
window [{}, {}): bytes before the "
+ "window were already consumed from the underlying stream
(totalBytesRead={})",
+ target,
+ windowStart,
+ windowEnd,
+ totalBytesRead_);
+ ranges_[0].position = static_cast<int32_t>(target - windowStart);
+}
+
+RowVectorPtr VeloxRssSortShuffleReaderDeserializer::readPage() {
+ using facebook::velox::serializer::presto::detail::kCompressedBitMask;
+ using facebook::velox::serializer::presto::detail::kHeaderSize;
+ using facebook::velox::serializer::presto::detail::PrestoHeader;
+ constexpr int32_t kPrestoHeaderSize = kHeaderSize;
+
+ // Fast path: peek the header without consuming; if the whole page fits in
+ // the current window, deserialize in-situ from in_ directly.
+ if (in_->remainingInWindow() >= kPrestoHeaderSize) {
+ std::string_view window(reinterpret_cast<const char*>(in_->data()),
in_->remainingInWindow());
+ auto peekedHeader = PrestoHeader::read(&window);
+ if (peekedHeader.has_value()) {
+ const int32_t payloadSize = (peekedHeader->pageCodecMarker &
kCompressedBitMask) != 0
+ ? peekedHeader->compressedSize
+ : peekedHeader->uncompressedSize;
+ const int64_t totalSize = kPrestoHeaderSize +
static_cast<int64_t>(payloadSize);
+ if (totalSize <= in_->remainingInWindow()) {
+ RowVectorPtr rowVector;
+ VectorStreamGroup::read(
+ in_.get(), memoryManager_->getLeafMemoryPool().get(), rowType_,
serde_, &rowVector, &serdeOptions_);
+ return rowVector;
+ }
+ }
+ }
+
+ // Slow path: the page spans multiple read windows. Reassemble it into a
+ // contiguous BufferInputStream so the serde's backward seek never touches
+ // window data overwritten by a refill.
+ std::array<uint8_t, kHeaderSize> headerStorage;
+ in_->readBytes(headerStorage.data(), kHeaderSize);
+ std::string_view headerBytes(reinterpret_cast<const
char*>(headerStorage.data()), kHeaderSize);
+ auto headerOpt = PrestoHeader::read(&headerBytes);
+ VELOX_CHECK(headerOpt.has_value(), "Invalid Presto page header");
+ const auto& header = *headerOpt;
+
+ const int32_t payloadSize =
+ (header.pageCodecMarker & kCompressedBitMask) != 0 ?
header.compressedSize : header.uncompressedSize;
+
+ // Payload is still in the window: stitch [copied header, in-situ payload].
+ if (payloadSize <= in_->remainingInWindow()) {
+ in_->advance(payloadSize);
+ BufferInputStream pageStream(std::vector<ByteRange>{
+ ByteRange{headerStorage.data(), kPrestoHeaderSize, 0},
ByteRange{in_->data() - payloadSize, payloadSize, 0}});
+ RowVectorPtr rowVector;
+ VectorStreamGroup::read(
+ &pageStream, memoryManager_->getLeafMemoryPool().get(), rowType_,
serde_, &rowVector, &serdeOptions_);
+ return rowVector;
+ }
+
+ // Payload spans windows: copy it into a contiguous buffer.
+ auto payloadBuffer = AlignedBuffer::allocate<char>(payloadSize,
memoryManager_->getLeafMemoryPool().get());
+ in_->readBytes(payloadBuffer->asMutable<uint8_t>(), payloadSize);
+ BufferInputStream pageStream(std::vector<ByteRange>{
+ ByteRange{headerStorage.data(), kPrestoHeaderSize, 0},
+ ByteRange{payloadBuffer->asMutable<uint8_t>(), payloadSize, 0}});
+ RowVectorPtr rowVector;
+ VectorStreamGroup::read(
+ &pageStream, memoryManager_->getLeafMemoryPool().get(), rowType_,
serde_, &rowVector, &serdeOptions_);
+ return rowVector;
+}
+
VeloxShuffleReader::VeloxShuffleReader(
const std::shared_ptr<arrow::Schema>& schema,
VeloxMemoryManager* memoryManager,
diff --git a/cpp/velox/shuffle/VeloxShuffleReader.h
b/cpp/velox/shuffle/VeloxShuffleReader.h
index 5daa0c0067..e35bd24f00 100644
--- a/cpp/velox/shuffle/VeloxShuffleReader.h
+++ b/cpp/velox/shuffle/VeloxShuffleReader.h
@@ -176,10 +176,12 @@ class VeloxRssSortShuffleReaderDeserializer : public
ShuffleReaderDeserializer {
std::unique_ptr<ColumnarBatchIterator> deserializeStreams() override;
private:
- class VeloxInputStream;
+ class RssSortShuffleReaderInputStream;
void loadNextStream();
+ facebook::velox::RowVectorPtr readPage();
+
std::shared_ptr<StreamReader> streamReader_;
VeloxMemoryManager* memoryManager_;
facebook::velox::RowTypePtr rowType_;
@@ -189,7 +191,7 @@ class VeloxRssSortShuffleReaderDeserializer : public
ShuffleReaderDeserializer {
facebook::velox::VectorSerde* const serde_;
facebook::velox::serializer::presto::PrestoVectorSerde::PrestoOptions
serdeOptions_;
int64_t& deserializeTime_;
- std::shared_ptr<VeloxInputStream> in_{nullptr};
+ std::shared_ptr<RssSortShuffleReaderInputStream> in_{nullptr};
std::shared_ptr<arrow::io::InputStream> arrowIn_{nullptr};
bool reachedEos_{false};
diff --git a/cpp/velox/tests/VeloxShuffleReaderTest.cc
b/cpp/velox/tests/VeloxShuffleReaderTest.cc
index 7cfab92c78..3273e9bb41 100644
--- a/cpp/velox/tests/VeloxShuffleReaderTest.cc
+++ b/cpp/velox/tests/VeloxShuffleReaderTest.cc
@@ -23,6 +23,9 @@
// - graceful EOS on an empty stream (e.g. an empty Celeborn partition);
// - EOS hit mid-page on a truncated compressed page;
// - a buggy upstream whose Read() returns an error status;
+// - uncompressed Presto pages spanning multiple read windows (nested struct
+// pre-scan, header crossing a refill boundary, checksummed pages) and the
+// single-window zero-copy fast path;
//
#include <gtest/gtest.h>
@@ -34,12 +37,14 @@
#include <cstdint>
#include <cstring>
#include <memory>
+#include <sstream>
#include <string>
#include <unordered_map>
#include <vector>
#include "compute/VeloxBackend.h"
#include "config/GlutenConfig.h"
+#include "memory/VeloxColumnarBatch.h"
#include "memory/VeloxMemoryManager.h"
#include "shuffle/VeloxShuffleReader.h"
#include "tests/utils/TestAllocationListener.h"
@@ -47,6 +52,7 @@
#include "velox/common/base/tests/GTestUtils.h"
#include "velox/serializers/PrestoSerializer.h"
#include "velox/type/Type.h"
+#include "velox/vector/VectorStream.h"
#include "velox/vector/tests/utils/VectorTestBase.h"
using namespace facebook::velox;
@@ -58,6 +64,9 @@ namespace {
// A minimal arrow::io::InputStream backed by a fixed in-memory payload. Once
// the payload is exhausted, Read returns 0 (EOS). With `errorRead`, Read
// returns an IOError instead, modeling a buggy upstream that fails the read.
+// With `firstReadLimit` >= 0, only the FIRST Read is capped to that many bytes
+// (subsequent reads are unbounded), modeling an upstream (e.g. a network
+// stream) whose first chunk ends mid-page.
//
// To keep a possible reader-side infinite loop (readBytes -> next() -> EOS ->
// silently return -> spin) from hanging the test until the CI timeout, Read
@@ -65,8 +74,8 @@ namespace {
// readers probe EOS only a couple of times, so the cap never trips for them.
class FakeInputStream final : public arrow::io::InputStream {
public:
- explicit FakeInputStream(std::vector<uint8_t> payload = {}, bool errorRead =
false)
- : payload_(std::move(payload)), errorRead_(errorRead) {}
+ explicit FakeInputStream(std::vector<uint8_t> payload = {}, bool errorRead =
false, int64_t firstReadLimit = -1)
+ : payload_(std::move(payload)), errorRead_(errorRead),
firstReadLimit_(firstReadLimit) {}
arrow::Status Close() override {
closed_ = true;
@@ -84,6 +93,10 @@ class FakeInputStream final : public arrow::io::InputStream {
return arrow::Status::IOError("fake upstream read failure");
}
int64_t toRead = std::min<int64_t>(nbytes,
static_cast<int64_t>(payload_.size()) - pos_);
+ if (firstRead_ && firstReadLimit_ >= 0) {
+ toRead = std::min<int64_t>(toRead, firstReadLimit_);
+ firstRead_ = false;
+ }
if (toRead > 0) {
std::memcpy(out, payload_.data() + pos_, toRead);
pos_ += toRead;
@@ -113,6 +126,8 @@ class FakeInputStream final : public arrow::io::InputStream
{
std::vector<uint8_t> payload_;
int64_t pos_{0};
bool errorRead_{false};
+ int64_t firstReadLimit_{-1};
+ bool firstRead_{true};
int32_t consecutiveEosReads_{0};
bool closed_{false};
};
@@ -129,7 +144,7 @@ void appendLe(std::vector<uint8_t>& out, T value) {
// Build a truncated Presto compressed page: a valid 21-byte header declaring
// compressedSize bytes of body, but only `bodyBytes` bytes follow. The
// reader's compressed branch calls source->readBytes(buf, compressedSize);
-// when EOS is hit mid-drain, GlutenByteInputStream::readBytes loops to
+// when EOS is hit mid-drain, RssSortShuffleReaderInputStream::readBytes loops
to
// next(true) which must VELOX_FAIL instead of spinning.
//
// Header layout (PrestoHeader.cpp): numRows:int32, pageCodecMarker:int8,
@@ -166,17 +181,59 @@ class VeloxShuffleReaderTest : public ::testing::Test,
public test::VectorTestBa
VeloxBackend::get()->tearDown();
}
- std::shared_ptr<VeloxRssSortShuffleReaderDeserializer>
makeDeserializer(std::shared_ptr<arrow::io::InputStream> in) {
+ std::shared_ptr<VeloxRssSortShuffleReaderDeserializer> makeDeserializer(
+ std::shared_ptr<arrow::io::InputStream> in,
+ const RowTypePtr& rowType = ROW({"c0"}, {INTEGER()})) {
auto streamReader = std::make_shared<TestStreamReader>(std::move(in));
return std::make_shared<VeloxRssSortShuffleReaderDeserializer>(
streamReader,
getDefaultMemoryManager(),
- ROW({"c0"}, {INTEGER()}),
+ rowType,
/*batchSize=*/1024,
common::CompressionKind_NONE,
deserializeTime_);
}
+ // Serializes `rowVector` into a single uncompressed Presto page
+ // (21-byte header + payload), the exact wire format the rss-sort writer
+ // produces. With `withChecksum`, a PrestoOutputStreamListener is attached so
+ // the writer fills in the checksum bit and CRC (same mechanism as
+ // VeloxHashShuffleWriter's complex-type flush). NOTE: must not use
+ // gluten::BufferOutputStream here — its write() ignores the listener.
+ std::vector<uint8_t> serializePage(const RowVectorPtr& rowVector, bool
withChecksum = false) {
+ serializer::presto::PrestoVectorSerde::PrestoOptions options;
+ options.compressionKind = common::CompressionKind_NONE;
+ auto serde = std::make_unique<serializer::presto::PrestoVectorSerde>();
+ VectorStreamGroup group(pool(), serde.get());
+ group.createStreamTree(asRowType(rowVector->type()), rowVector->size(),
&options);
+ group.append(rowVector);
+ serializer::presto::PrestoOutputStreamListener listener;
+ std::stringstream out;
+ facebook::velox::OStreamOutputStream os(&out, withChecksum ? &listener :
nullptr);
+ group.flush(&os);
+ const auto str = out.str();
+ return std::vector<uint8_t>(str.begin(), str.end());
+ }
+
+ // ROW<c0: ARRAY<ROW<a: INTEGER>>> with `numArrays` arrays of
+ // `elementsPerArray` dense elements each. The nested struct triggers the
+ // Presto serde's pre-scan (tellp -> scan page -> seekp back).
+ RowVectorPtr makeNestedArraysRowVector(vector_size_t numArrays,
vector_size_t elementsPerArray) {
+ const vector_size_t numElements = numArrays * elementsPerArray;
+ auto offsets = AlignedBuffer::allocate<vector_size_t>(numArrays, pool());
+ auto sizes = AlignedBuffer::allocate<vector_size_t>(numArrays, pool());
+ auto* rawOffsets = offsets->asMutable<vector_size_t>();
+ auto* rawSizes = sizes->asMutable<vector_size_t>();
+ for (vector_size_t i = 0; i < numArrays; ++i) {
+ rawOffsets[i] = i * elementsPerArray;
+ rawSizes[i] = elementsPerArray;
+ }
+ auto elements = makeRowVector({makeFlatVector<int32_t>(numElements,
[](vector_size_t row) { return row % 1024; })});
+ auto arrayVector = std::make_shared<ArrayVector>(
+ pool(), ARRAY(ROW({"a"}, {INTEGER()})), BufferPtr(nullptr), numArrays,
offsets, sizes, elements);
+ return makeRowVector({arrayVector});
+ }
+
int64_t deserializeTime_{0};
};
@@ -196,8 +253,7 @@ TEST_F(VeloxShuffleReaderTest, EosMidPageThrows) {
auto payload = buildTruncatedCompressedPage(/*compressedSize=*/1000,
/*bodyBytes=*/8);
auto deserializer =
makeDeserializer(std::make_shared<FakeInputStream>(std::move(payload)));
- VELOX_ASSERT_THROW(
- deserializer->next(), "Reading past end of
VeloxRssSortShuffleReaderDeserializer::VeloxInputStream");
+ VELOX_ASSERT_THROW(deserializer->next(), "Reading past end of
RssSortShuffleReaderInputStream");
}
// A buggy upstream whose Read returns an error status. The reader must
@@ -208,4 +264,98 @@ TEST_F(VeloxShuffleReaderTest, ErrorReadThrows) {
EXPECT_THROW((void)deserializer->next(), GlutenException);
}
+// Multi-window bug regression: uncompressed + nested struct + page spanning
+// multiple read windows is the exact combination that reproduced the original
+// failure (corrupted data / spurious "Reading past end" EOS). With the fix the
+// page deserializes correctly.
+TEST_F(VeloxShuffleReaderTest, UncompressedNestedStructPageSpansWindows) {
+ constexpr vector_size_t kNumArrays = 20000;
+ constexpr vector_size_t kElementsPerArray = 30;
+
+ auto rowVector = makeNestedArraysRowVector(kNumArrays, kElementsPerArray);
+ auto payload = serializePage(rowVector);
+ // The page must be larger than the reader's read window (~1MB) to exercise
+ // the multi-window slow path.
+ ASSERT_GT(payload.size(), 1 << 20);
+
+ auto deserializer =
+ makeDeserializer(std::make_shared<FakeInputStream>(std::move(payload)),
asRowType(rowVector->type()));
+ auto batch = deserializer->next();
+ ASSERT_NE(batch, nullptr);
+ auto result = VeloxColumnarBatch::from(pool(), batch)->getRowVector();
+ assertEqualVectors(rowVector, result);
+ ASSERT_EQ(deserializer->next(), nullptr);
+}
+
+// Fast path: a small page that fits in one window is deserialized in-situ
+// from the read window (zero copy).
+TEST_F(VeloxShuffleReaderTest, SingleWindowPageZeroCopy) {
+ constexpr vector_size_t kNumRows = 200;
+ auto rowVector = makeRowVector(
+ {makeFlatVector<int32_t>(kNumRows, [](vector_size_t row) { return
static_cast<int32_t>(row * 7); })});
+
+ auto payload = serializePage(rowVector);
+ // Small page: header + payload well under 1MB -> zero-copy fast path.
+ ASSERT_LT(payload.size(), 1 << 20);
+
+ auto deserializer =
+ makeDeserializer(std::make_shared<FakeInputStream>(std::move(payload)),
asRowType(rowVector->type()));
+ auto batch = deserializer->next();
+ ASSERT_NE(batch, nullptr);
+ auto result = VeloxColumnarBatch::from(pool(), batch)->getRowVector();
+ assertEqualVectors(rowVector, result);
+ ASSERT_EQ(deserializer->next(), nullptr);
+}
+
+// Slow path 2: the page header crosses a refill boundary. readPage() copies
+// the header out (readBytes refills mid-header) and stitches it with the
+// in-situ payload into a contiguous stream before deserializing.
+TEST_F(VeloxShuffleReaderTest, HeaderCrossesRefillBoundary) {
+ constexpr vector_size_t kArraysA = 1024;
+ constexpr vector_size_t kArraysB = 100;
+ constexpr vector_size_t kElementsPerArray = 30;
+
+ auto rowVectorA = makeNestedArraysRowVector(kArraysA, kElementsPerArray);
+ auto rowVectorB = makeNestedArraysRowVector(kArraysB, kElementsPerArray);
+ auto pageA = serializePage(rowVectorA);
+ auto pageB = serializePage(rowVectorB);
+
+ std::vector<uint8_t> payload = pageA;
+ payload.insert(payload.end(), pageB.begin(), pageB.end());
+ // First window holds all of A plus only 10 bytes of B's header.
+ const int64_t firstReadLimit = static_cast<int64_t>(pageA.size()) + 10;
+
+ auto deserializer = makeDeserializer(
+ std::make_shared<FakeInputStream>(std::move(payload),
/*errorRead=*/false, firstReadLimit),
+ asRowType(rowVectorA->type()));
+
+ // Page A fits the window -> zero-copy fast path.
+ auto batchA = deserializer->next();
+ ASSERT_NE(batchA, nullptr);
+ assertEqualVectors(rowVectorA, VeloxColumnarBatch::from(pool(),
batchA)->getRowVector());
+
+ // Page B's header crossed the refill boundary -> stitched mid path.
+ auto batchB = deserializer->next();
+ ASSERT_NE(batchB, nullptr);
+ assertEqualVectors(rowVectorB, VeloxColumnarBatch::from(pool(),
batchB)->getRowVector());
+
+ ASSERT_EQ(deserializer->next(), nullptr);
+}
+
+// Checksummed page: the serde scans the payload via nextView() and seeks
+// back to verify the CRC before deserializing.
+TEST_F(VeloxShuffleReaderTest, PageDeserializesWithChecksum) {
+ auto rowVector =
+ makeRowVector({makeFlatVector<int32_t>(200, [](vector_size_t row) {
return static_cast<int32_t>(row * 7); })});
+ auto payload = serializePage(rowVector, /*withChecksum=*/true);
+
+ auto deserializer =
+ makeDeserializer(std::make_shared<FakeInputStream>(std::move(payload)),
asRowType(rowVector->type()));
+ auto batch = deserializer->next();
+ ASSERT_NE(batch, nullptr);
+ auto result = VeloxColumnarBatch::from(pool(), batch)->getRowVector();
+ assertEqualVectors(rowVector, result);
+ ASSERT_EQ(deserializer->next(), nullptr);
+}
+
} // namespace gluten
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]