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_ = &range;
-        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]

Reply via email to