#include <var_opt_sketch.hpp>
#include <catch2/catch.hpp>
#include <vector>
#include <string>
#include <sstream>
#include <fstream>
#include <cmath>
#include <random>
#include <stdexcept>
#ifdef TEST_BINARY_INPUT_PATH
static std::string testBinaryInputPath = TEST_BINARY_INPUT_PATH;
#else
static std::string testBinaryInputPath = "test/";
#endif
namespace datasketches {
static constexpr double EPS = 1e-13;
static var_opt_sketch<int> create_unweighted_sketch(uint32_t k, uint64_t n) {
var_opt_sketch<int> sk(k);
for (uint64_t i = 0; i < n; ++i) {
sk.update(static_cast<int>(i), 1.0);
}
return sk;
}
template<typename T, typename A>
static void check_if_equal(var_opt_sketch<T, A>& sk1, var_opt_sketch<T, A>& sk2) {
REQUIRE(sk1.get_k() == sk2.get_k());
REQUIRE(sk1.get_n() == sk2.get_n());
REQUIRE(sk1.get_num_samples() == sk2.get_num_samples());
auto it1 = sk1.begin();
auto it2 = sk2.begin();
while ((it1 != sk1.end()) && (it2 != sk2.end())) {
auto p1 = *it1;
auto p2 = *it2;
REQUIRE(p1.first == p2.first); REQUIRE(p1.second == p2.second); ++it1;
++it2;
}
REQUIRE((it1 == sk1.end() && it2 == sk2.end())); }
TEST_CASE("varopt sketch: invalid k", "[var_opt_sketch]") {
REQUIRE_THROWS_AS(var_opt_sketch<int>(0), std::invalid_argument);
REQUIRE_THROWS_AS(var_opt_sketch<int>(1U << 31), std::invalid_argument); }
TEST_CASE("varopt sketch: bad serialization version", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(16, 16);
std::vector<uint8_t> bytes = sk.serialize();
bytes[1] = 0;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
std::stringstream ss;
std::string str(bytes.begin(), bytes.end());
ss.str(str);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::invalid_argument);
}
TEST_CASE("varopt sketch: bad family", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(16, 16);
std::vector<uint8_t> bytes = sk.serialize();
bytes[2] = 0;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
std::stringstream ss;
std::string str(bytes.begin(), bytes.end());
ss.str(str);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::invalid_argument);
}
TEST_CASE("varopt sketch: bad prelongs", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(32, 33);
std::vector<uint8_t> bytes = sk.serialize();
bytes[0] = 0; REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
bytes[0] = 2; REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
bytes[0] = 5; REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
}
TEST_CASE("varopt sketch: malformed preamble", "[var_opt_sketch]") {
uint32_t k = 50;
var_opt_sketch<int> sk = create_unweighted_sketch(k, k);
const std::vector<uint8_t> src_bytes = sk.serialize();
std::vector<uint8_t> bytes(src_bytes);
bytes[0] = 4; REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
bytes = src_bytes;
*reinterpret_cast<int32_t*>(&bytes[4]) = 0;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
bytes = src_bytes;
*reinterpret_cast<int32_t*>(&bytes[16]) = -1;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
bytes = src_bytes;
*reinterpret_cast<int32_t*>(&bytes[20]) = -128;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
}
TEST_CASE("varopt sketch: empty sketch", "[var_opt_sketch]") {
var_opt_sketch<std::string> sk(5);
REQUIRE(sk.get_n() == 0);
REQUIRE(sk.get_num_samples() == 0);
std::vector<uint8_t> bytes = sk.serialize();
REQUIRE(bytes.size() == (1 << 3));
var_opt_sketch<std::string> loaded_sk = var_opt_sketch<std::string>::deserialize(bytes.data(), bytes.size());
REQUIRE(loaded_sk.get_n() == 0);
REQUIRE(loaded_sk.get_num_samples() == 0);
}
TEST_CASE("varopt sketch: non-empty degenerate sketch", "[var_opt_sketch]") {
var_opt_sketch<std::string> sk(12, resize_factor::X2);
std::vector<uint8_t> bytes = sk.serialize();
while (bytes.size() < 24) { bytes.push_back((uint8_t) 0);
}
bytes[3] = 0;
REQUIRE_THROWS_AS(var_opt_sketch<std::string>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
}
TEST_CASE("varopt sketch: invalid weight", "[var_opt_sketch]") {
var_opt_sketch<std::string> sk(100, resize_factor::X2);
REQUIRE_THROWS_AS(sk.update("invalid_weight", -1.0), std::invalid_argument);
sk.update("zero weight", 0.0);
REQUIRE(sk.is_empty());
}
TEST_CASE("varopt sketch: corrupt serialized weight", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(100, 20);
auto bytes = sk.serialize();
size_t preamble_bytes = (bytes[0] & 0x3f) << 3;
*reinterpret_cast<double*>(&bytes[preamble_bytes]) = -1.5;
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
std::stringstream ss(std::ios::in | std::ios::out | std::ios::binary);
for (auto& b : bytes) { ss >> b; }
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::invalid_argument);
}
TEST_CASE("varopt sketch: cumulative weight", "[var_opt_sketch]") {
uint32_t k = 256;
uint64_t n = 10 * k;
var_opt_sketch<int> sk(k);
std::random_device rd; std::mt19937_64 rand(rd());
std::normal_distribution<double> N(0.0, 1.0);
double input_sum = 0.0;
for (size_t i = 0; i < n; ++i) {
double w = std::exp(5 * N(rand));
input_sum += w;
sk.update(static_cast<int>(i), w);
}
double output_sum = 0.0;
for (auto pair : sk) { output_sum += pair.second;
}
double weight_ratio = output_sum / input_sum;
REQUIRE(weight_ratio == Approx(1.0).margin(EPS));
}
TEST_CASE("varopt sketch: under-full sketch serialization", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(100, 10);
auto bytes = sk.serialize();
var_opt_sketch<int> sk_from_bytes = var_opt_sketch<int>::deserialize(bytes.data(), bytes.size());
check_if_equal(sk, sk_from_bytes);
std::stringstream ss(std::ios::in | std::ios::out | std::ios::binary);
sk.serialize(ss);
var_opt_sketch<int> sk_from_stream = var_opt_sketch<int>::deserialize(ss);
check_if_equal(sk, sk_from_stream);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size() - 1), std::out_of_range);
std::string str_trunc((char*)&bytes[0], bytes.size() - 1);
ss.str(str_trunc);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::runtime_error);
}
TEST_CASE("varopt sketch: end-of-warmup sketch serialization", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(2843, 2843); auto bytes = sk.serialize();
REQUIRE((bytes.data()[0] & 0x3f) == 3);
var_opt_sketch<int> sk_from_bytes = var_opt_sketch<int>::deserialize(bytes.data(), bytes.size());
check_if_equal(sk, sk_from_bytes);
std::stringstream ss(std::ios::in | std::ios::out | std::ios::binary);
sk.serialize(ss);
var_opt_sketch<int> sk_from_stream = var_opt_sketch<int>::deserialize(ss);
check_if_equal(sk, sk_from_stream);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size() - 1000), std::out_of_range);
std::string str_trunc((char*)&bytes[0], bytes.size() - 100);
ss.str(str_trunc);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::runtime_error);
}
TEST_CASE("varopt sketch: full sketch serialization", "[var_opt_sketch]") {
var_opt_sketch<int> sk = create_unweighted_sketch(32, 32);
sk.update(100, 100.0);
sk.update(101, 101.0);
subset_summary summary = sk.estimate_subset_sum([](int){ return true; });
double total_weight = summary.total_sketch_weight;
double cum_weight = 0.0;
for (auto pair : sk) {
cum_weight += pair.second;
}
double weight_ratio = cum_weight / total_weight;
REQUIRE(weight_ratio == Approx(1.0).margin(EPS));
auto it = sk.begin();
auto p1 = *it;
++it;
auto p2 = *it;
REQUIRE(p1.second == Approx(100.0).margin(EPS));
REQUIRE(p2.second == Approx(101.0).margin(EPS));
REQUIRE(p1.first == 100);
REQUIRE(p2.first == 101);
REQUIRE(it->first == p2.first);
REQUIRE(it->second == p2.second);
auto bytes = sk.serialize();
REQUIRE((bytes.data()[0] & 0x3f) == 4);;
auto sk_from_bytes = var_opt_sketch<int>::deserialize(bytes.data(), bytes.size());
check_if_equal(sk, sk_from_bytes);
std::stringstream ss(std::ios::in | std::ios::out | std::ios::binary);
sk.serialize(ss);
auto sk_from_stream = var_opt_sketch<int>::deserialize(ss);
check_if_equal(sk, sk_from_stream);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(bytes.data(), bytes.size() - 100), std::out_of_range);
std::string str_trunc((char*)&bytes[0], bytes.size() - 100);
ss.str(str_trunc);
REQUIRE_THROWS_AS(var_opt_sketch<int>::deserialize(ss), std::runtime_error);
}
TEST_CASE("varopt sketch: string serialization", "[var_opt_sketch]") {
var_opt_sketch<std::string> sk(5);
sk.update("a", 1.0);
sk.update("bc", 1.0);
sk.update("def", 1.0);
sk.update("ghij", 1.0);
sk.update("klmno", 1.0);
sk.update("heavy item", 100.0);
auto bytes = sk.serialize();
var_opt_sketch<std::string> sk_from_bytes = var_opt_sketch<std::string>::deserialize(bytes.data(), bytes.size());
check_if_equal(sk, sk_from_bytes);
std::stringstream ss(std::ios::in | std::ios::out | std::ios::binary);
sk.serialize(ss);
var_opt_sketch<std::string> sk_from_stream = var_opt_sketch<std::string>::deserialize(ss);
check_if_equal(sk, sk_from_stream);
REQUIRE_THROWS_AS(var_opt_sketch<std::string>::deserialize(bytes.data(), bytes.size() - 12), std::out_of_range);
std::string str_trunc((char*)&bytes[0], bytes.size() - 12);
ss.str(str_trunc);
REQUIRE_THROWS_AS(var_opt_sketch<std::string>::deserialize(ss), std::runtime_error);
}
TEST_CASE("varopt sketch: pseudo-light update", "[var_opt_sketch]") {
uint32_t k = 1024;
var_opt_sketch<int> sk = create_unweighted_sketch(k, k + 1);
sk.update(0, 1.0);
auto it = sk.begin();
double wt = (*it).second;
REQUIRE(wt == Approx((k + 2.0) / k).margin(EPS));
subset_summary summary = sk.estimate_subset_sum([](int){ return true; });
double total_weight = summary.total_sketch_weight;
double cum_weight = 0.0;
for (auto pair : sk) {
cum_weight += pair.second;
}
double weight_ratio = cum_weight / total_weight;
REQUIRE(weight_ratio == Approx(1.0).margin(EPS));
}
TEST_CASE("varopt sketch: pseudo-heavy update", "[var_opt_sketch]") {
uint32_t k = 1024;
double wt_scale = 10.0 * k;
var_opt_sketch<int> sk = create_unweighted_sketch(k, k + 1);
for (uint32_t i = 1; i <= k; ++i) {
sk.update(-1 * static_cast<int>(i), k + (i * wt_scale));
}
auto it = sk.begin();
double wt = (*it).second;
REQUIRE(wt == Approx(1.0 * (k + (2 * wt_scale))).margin(EPS));
while (it != sk.end()) {
wt = (*it).second;
++it;
}
REQUIRE(wt == Approx(1.0 + wt_scale + (2 * k)).margin(EPS));
}
TEST_CASE("varopt sketch: reset", "[var_opt_sketch]") {
uint32_t k = 1024;
uint64_t n1 = 20;
uint64_t n2 = 2 * k;
var_opt_sketch<std::string> sk(k);
for (uint64_t i = 0; i < n2; ++i) {
sk.update(std::to_string(i), 100.0 + i);
}
REQUIRE(sk.get_n() == n2);
REQUIRE(sk.get_k() == k);
sk.reset();
REQUIRE(sk.get_n() == 0);
REQUIRE(sk.get_k() == k);
for (uint64_t i = 0; i < n1; ++i)
sk.update(std::to_string(i));
REQUIRE(sk.get_n() == n1);
REQUIRE(sk.get_k() == k);
sk.reset();
REQUIRE(sk.get_n() == 0);
REQUIRE(sk.get_k() == k);
}
TEST_CASE("varopt sketch: estimate subset sum", "[var_opt_sketch]") {
uint32_t k = 10;
var_opt_sketch<int> sk(k);
subset_summary summary = sk.estimate_subset_sum([](int){ return true; });
REQUIRE(summary.estimate == 0.0);
REQUIRE(summary.total_sketch_weight == 0.0);
double total_weight = 0.0;
for (uint32_t i = 1; i <= (k - 1); ++i) {
sk.update(i, 1.0 * i);
total_weight += 1.0 * i;
}
summary = sk.estimate_subset_sum([](int){ return true; });
REQUIRE(summary.estimate == total_weight);
REQUIRE(summary.lower_bound == total_weight);
REQUIRE(summary.upper_bound == total_weight);
REQUIRE(summary.total_sketch_weight == total_weight);
for (uint32_t i = k; i <= (k + 1); ++i) {
sk.update(i, 1.0 * i);
total_weight += 1.0 * i;
}
summary = sk.estimate_subset_sum([](int){ return true; });
REQUIRE(summary.estimate == Approx(total_weight).margin(EPS));
REQUIRE(summary.upper_bound == Approx(total_weight).margin(EPS));
REQUIRE(summary.lower_bound < total_weight);
REQUIRE(summary.total_sketch_weight == Approx(total_weight).margin(EPS));
summary = sk.estimate_subset_sum([](int){ return false; });
REQUIRE(summary.estimate == 0.0);
REQUIRE(summary.lower_bound == 0.0);
REQUIRE(summary.upper_bound > 0.0);
REQUIRE(summary.total_sketch_weight == Approx(total_weight).margin(EPS));
for (uint32_t i = 1; i <= (k + 1); ++i) {
sk.update(-1 * static_cast<int32_t>(i), static_cast<double>(i));
total_weight += 1.0 * i;
}
summary = sk.estimate_subset_sum([](int x) { return x < 0; });
REQUIRE(summary.estimate >= summary.lower_bound);
REQUIRE(summary.estimate <= summary.upper_bound);
REQUIRE(summary.lower_bound < (total_weight / 1.4));
REQUIRE(summary.upper_bound > (total_weight / 2.6));
REQUIRE(summary.total_sketch_weight == Approx(total_weight).margin(EPS));
var_opt_sketch<bool> sk2(k);
total_weight = 0.0;
for (uint32_t i = 1; i <= (k - 1); ++i) {
sk2.update((i % 2) == 0, 1.0 * i);
total_weight += i;
}
summary = sk2.estimate_subset_sum([](bool b){ return !b; });
REQUIRE(summary.estimate == summary.lower_bound);
REQUIRE(summary.estimate == summary.upper_bound);
REQUIRE(summary.estimate < total_weight); }
}