diff --git a/cpp/src/arrow/stl_allocator.h b/cpp/src/arrow/stl_allocator.h index bb52767282ec..62585809c785 100644 --- a/cpp/src/arrow/stl_allocator.h +++ b/cpp/src/arrow/stl_allocator.h @@ -19,6 +19,7 @@ #include #include +#include #include #include #include @@ -72,6 +73,9 @@ class allocator { const_pointer address(const_reference r) const noexcept { return std::addressof(r); } pointer allocate(size_type n, const void* /*hint*/ = NULLPTR) { + if (n > size_max() || n > std::numeric_limits::max() / sizeof(T)) { + throw BadAlloc(Status::OutOfMemory("Memory allocation size too large")); + } uint8_t* data; Status s = pool_->Allocate(n * sizeof(T), &data); if (!s.ok()) throw BadAlloc(std::move(s)); diff --git a/cpp/src/arrow/stl_test.cc b/cpp/src/arrow/stl_test.cc index ce5adf0c0e26..39916e1fc009 100644 --- a/cpp/src/arrow/stl_test.cc +++ b/cpp/src/arrow/stl_test.cc @@ -558,6 +558,20 @@ TEST(allocator, MemoryTracking) { ASSERT_EQ(0, pool->bytes_allocated()); } +TEST(allocator, AllocationSizeOverflow) { + allocator alloc; + const size_t first_overflow = std::numeric_limits::max() / sizeof(uint64_t) + 1; + for (const size_t n : {first_overflow, first_overflow + 1}) { + SCOPED_TRACE(n); + EXPECT_THROW( + { + auto* data = alloc.allocate(n); + alloc.deallocate(data, n); + }, + std::bad_alloc); + } +} + #if !(defined(ARROW_VALGRIND) || defined(ADDRESS_SANITIZER) || defined(ARROW_JEMALLOC)) TEST(allocator, TestOOM) {