From c761c5d83a38a29d91ff1cbc8e837c0a6f2aec6e Mon Sep 17 00:00:00 2001 From: Shani Elharrar Date: Wed, 7 Oct 2026 14:09:04 +0300 Subject: [PATCH 1/2] Fix infinite loop in kll_helper::floor_of_log2_of_fraction for large numerators Doubling denom until it exceeds numer overflows to zero once numer >= 2^63, so the loop never terminated. ub_on_num_levels(n) calls it with the merged stream weight, which made a merge reaching n >= 2^63 hang. Compare against numer / 2 instead, which is equivalent and cannot overflow. Co-Authored-By: Claude Opus 5.5 --- kll/include/kll_helper_impl.hpp | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/kll/include/kll_helper_impl.hpp b/kll/include/kll_helper_impl.hpp index 31534d9a..56b5be69 100644 --- a/kll/include/kll_helper_impl.hpp +++ b/kll/include/kll_helper_impl.hpp @@ -39,8 +39,9 @@ uint8_t kll_helper::floor_of_log2_of_fraction(uint64_t numer, uint64_t denom) { if (denom > numer) { return 0; } uint8_t count = 0; while (true) { + // denom * 2 > numer, without overflowing denom + if (denom > (numer >> 1)) { return count; } denom <<= 1; - if (denom > numer) { return count; } count++; } } From f2728bfcedfad363a31b129191ca15adfd96e6be Mon Sep 17 00:00:00 2001 From: Shani Elharrar Date: Wed, 7 Oct 2026 14:09:04 +0300 Subject: [PATCH 2/2] Add weighted update to kll_sketch update(item, weight) is equivalent to calling update(item) weight times. A weight that fits into the free space of level zero is applied as plain updates; a larger one builds an exact sketch holding one copy of the item at each level whose bit is set in the weight and merges it, so the cost is logarithmic in the weight. This mirrors the weighted update in datasketches-java. Co-Authored-By: Claude Opus 5.5 --- kll/include/kll_sketch.hpp | 15 ++++ kll/include/kll_sketch_impl.hpp | 45 ++++++++++++ kll/test/kll_sketch_test.cpp | 120 ++++++++++++++++++++++++++++++++ 3 files changed, 180 insertions(+) diff --git a/kll/include/kll_sketch.hpp b/kll/include/kll_sketch.hpp index d672c419..0133b899 100644 --- a/kll/include/kll_sketch.hpp +++ b/kll/include/kll_sketch.hpp @@ -237,6 +237,18 @@ class kll_sketch { template void update(FwdT&& item); + /** + * Updates this sketch with the given data item repeated the given number of times. + * The result is equivalent to calling update(item) weight times, at a cost that grows + * with the logarithm of the weight rather than with the weight itself. + * If cross-language portability is required, callers should ensure that + * the input string uses a compatible encoding (valid UTF-8). + * @param item from a stream of items + * @param weight number of times the item is repeated, must be positive + */ + template + void update(FwdT&& item, uint64_t weight); + /** * Merges another sketch into this one. * If sketches contain strings, callers are responsible for ensuring that @@ -570,6 +582,9 @@ class kll_sketch { std::unique_ptr items, uint32_t items_size, optional&& min_item, optional&& max_item, bool is_level_zero_sorted, const C& comparator); + // for weighted update + kll_sketch(uint16_t k, const T& item, uint64_t weight, const C& comparator, const A& allocator); + // common update code inline void update_min_max(const T& item); inline uint32_t internal_update(); diff --git a/kll/include/kll_sketch_impl.hpp b/kll/include/kll_sketch_impl.hpp index b12a39c8..fcc51d3a 100644 --- a/kll/include/kll_sketch_impl.hpp +++ b/kll/include/kll_sketch_impl.hpp @@ -186,6 +186,51 @@ void kll_sketch::update(FwdT&& item) { reset_sorted_view(); } +template +template +void kll_sketch::update(FwdT&& item, uint64_t weight) { + if (weight == 0) throw std::invalid_argument("weight must be positive"); + if (!check_update_item(item)) { return; } + if (weight < levels_[0]) { // fits into level zero without compaction + update_min_max(static_cast(item)); + for (uint64_t i = 1; i < weight; ++i) new (&items_[internal_update()]) T(static_cast(item)); + new (&items_[internal_update()]) T(std::forward(item)); + reset_sorted_view(); + } else { + merge(kll_sketch(k_, static_cast(item), weight, comparator_, allocator_)); + } +} + +// An item at level h carries weight 2^h, so one copy of the item at each level +// whose bit is set in the weight is an exact sketch of that many copies. +// Level capacities are defined up to 61 levels, so bits above the top level +// fold into weight >> 60 copies at level 60. +template +kll_sketch::kll_sketch(uint16_t k, const T& item, uint64_t weight, const C& comparator, const A& allocator): +comparator_(comparator), +allocator_(allocator), +k_(k), +m_(kll_constants::DEFAULT_M), +min_k_(k), +num_levels_(std::min(64 - count_leading_zeros_in_u64(weight), 61)), +is_level_zero_sorted_(true), +n_(weight), +levels_(num_levels_ + 1, 0, allocator), +items_(nullptr), +items_size_(0), +min_item_(item), +max_item_(item), +sorted_view_(nullptr) +{ + for (uint8_t level = 0; level < num_levels_; ++level) { + const uint64_t count = level + 1 < num_levels_ ? (weight >> level) & 1 : weight >> level; + levels_[level + 1] = levels_[level] + static_cast(count); + } + items_size_ = levels_[num_levels_]; + items_ = allocator_.allocate(items_size_); + for (uint32_t i = 0; i < items_size_; ++i) new (&items_[i]) T(item); +} + template void kll_sketch::update_min_max(const T& item) { if (is_empty()) { diff --git a/kll/test/kll_sketch_test.cpp b/kll/test/kll_sketch_test.cpp index b8d0b031..d4aa396e 100644 --- a/kll/test/kll_sketch_test.cpp +++ b/kll/test/kll_sketch_test.cpp @@ -443,6 +443,9 @@ TEST_CASE("kll sketch", "[kll_sketch]") { REQUIRE(kll_helper::floor_of_log2_of_fraction(6, 2) == 1); REQUIRE(kll_helper::floor_of_log2_of_fraction(7, 2) == 1); REQUIRE(kll_helper::floor_of_log2_of_fraction(8, 2) == 2); + REQUIRE(kll_helper::floor_of_log2_of_fraction(1ULL << 63, 1) == 63); + REQUIRE(kll_helper::floor_of_log2_of_fraction(std::numeric_limits::max(), 1) == 63); + REQUIRE(kll_helper::floor_of_log2_of_fraction(std::numeric_limits::max(), 3) == 62); } SECTION("out of order split points, float") { @@ -831,6 +834,123 @@ TEST_CASE("kll sketch", "[kll_sketch]") { REQUIRE(sb.get_n() == 3); } + SECTION("weighted update: zero weight") { + kll_float_sketch sketch(200, std::less(), 0); + REQUIRE_THROWS_AS(sketch.update(1.0f, 0), std::invalid_argument); + REQUIRE(sketch.is_empty()); + } + + SECTION("weighted update: NaN") { + kll_float_sketch sketch(200, std::less(), 0); + sketch.update(std::numeric_limits::quiet_NaN(), 1000); + REQUIRE(sketch.is_empty()); + } + + SECTION("weighted update: same as repeated updates in exact mode") { + kll_float_sketch weighted(200, std::less(), 0); + kll_float_sketch repeated(200, std::less(), 0); + for (int i = 1; i <= 15; ++i) { + weighted.update(static_cast(i), i); + for (int j = 0; j < i; ++j) repeated.update(static_cast(i)); + } + REQUIRE(weighted.get_n() == repeated.get_n()); + REQUIRE(weighted.get_num_retained() == repeated.get_num_retained()); + REQUIRE_FALSE(weighted.is_estimation_mode()); + for (int i = 0; i <= 16; ++i) { + REQUIRE(weighted.get_rank(static_cast(i)) == repeated.get_rank(static_cast(i))); + REQUIRE(weighted.get_rank(static_cast(i), false) == repeated.get_rank(static_cast(i), false)); + } + } + + SECTION("weighted update: one item with large weight") { + kll_float_sketch sketch(200, std::less(), 0); + const uint64_t weight = (1ULL << 40) + 12345; + sketch.update(7.0f, weight); + REQUIRE(sketch.get_n() == weight); + REQUIRE(sketch.get_num_retained() == 7); // number of bits set in the weight + REQUIRE(sketch.get_min_item() == 7.0f); + REQUIRE(sketch.get_max_item() == 7.0f); + REQUIRE(sketch.get_rank(7.0f) == 1); + REQUIRE(sketch.get_rank(7.0f, false) == 0); + REQUIRE(sketch.get_quantile(0.5) == 7.0f); + + uint64_t total_weight = 0; + for (auto pair: sketch) total_weight += pair.second; + REQUIRE(total_weight == weight); + } + + SECTION("weighted update: maximum weight") { + kll_float_sketch sketch(200, std::less(), 0); + sketch.update(1.0f); + sketch.update(2.0f, std::numeric_limits::max() - 1); + REQUIRE(sketch.get_n() == std::numeric_limits::max()); + REQUIRE(sketch.get_min_item() == 1.0f); + REQUIRE(sketch.get_max_item() == 2.0f); + REQUIRE(sketch.get_quantile(0.5) == 2.0f); + + uint64_t total_weight = 0; + for (auto pair: sketch) total_weight += pair.second; + REQUIRE(total_weight == sketch.get_n()); + + auto bytes = sketch.serialize(); + auto sketch2 = kll_float_sketch::deserialize(bytes.data(), bytes.size(), serde(), std::less(), 0); + REQUIRE(sketch2.get_n() == sketch.get_n()); + REQUIRE(sketch2.get_num_retained() == sketch.get_num_retained()); + } + + SECTION("weighted update: large weights into a full sketch") { + kll_float_sketch sketch(200, std::less(), 0); + for (int i = 0; i < 1000; ++i) sketch.update(static_cast(i)); + sketch.update(-1.0f, 1000); + sketch.update(2000.0f, 3000); + REQUIRE(sketch.get_n() == 5000); + REQUIRE(sketch.get_min_item() == -1.0f); + REQUIRE(sketch.get_max_item() == 2000.0f); + REQUIRE(sketch.get_rank(-1.0f) == Approx(0.2).margin(RANK_EPS_FOR_K_200)); + REQUIRE(sketch.get_rank(2000.0f, false) == Approx(0.4).margin(RANK_EPS_FOR_K_200)); + } + + SECTION("weighted update: estimation mode") { + kll_float_sketch sketch(200, std::less(), 0); + const int n = 10000; + std::vector weights(n); + uint64_t total = 0; + for (int i = 0; i < n; ++i) { + weights[i] = static_cast(i % 13) * 1000 + 1; + total += weights[i]; + sketch.update(static_cast(i), weights[i]); + } + REQUIRE(sketch.get_n() == total); + REQUIRE(sketch.is_estimation_mode()); + REQUIRE(sketch.get_min_item() == 0); + REQUIRE(sketch.get_max_item() == n - 1); + + uint64_t weight_below = 0; + for (int i = 0; i < n; ++i) { + const double true_rank = static_cast(weight_below) / total; + REQUIRE(sketch.get_rank(static_cast(i), false) == Approx(true_rank).margin(RANK_EPS_FOR_K_200)); + weight_below += weights[i]; + } + + auto bytes = sketch.serialize(); + auto sketch2 = kll_float_sketch::deserialize(bytes.data(), bytes.size(), serde(), std::less(), 0); + REQUIRE(sketch2.get_n() == sketch.get_n()); + REQUIRE(sketch2.get_num_retained() == sketch.get_num_retained()); + REQUIRE(sketch2.get_quantile(0.5) == sketch.get_quantile(0.5)); + } + + SECTION("weighted update: strings") { + kll_string_sketch sketch(200, std::less(), 0); + sketch.update(std::string("b"), 3); + const std::string a("a"); + sketch.update(a, 1000); + sketch.update(std::string("c"), 1ULL << 20); + REQUIRE(sketch.get_n() == 3 + 1000 + (1ULL << 20)); + REQUIRE(sketch.get_min_item() == "a"); + REQUIRE(sketch.get_max_item() == "c"); + REQUIRE(sketch.get_rank("a") == Approx(1000.0 / sketch.get_n()).margin(RANK_EPS_FOR_K_200)); + } + // cleanup REQUIRE(test_allocator_total_bytes == 0); }