#include "array_of_doubles_sketch_shim.h"
#include "array_of_doubles_compact_shim.h"
#include "array_of_doubles_sketch.rs.h"
namespace apache_datasketches_rs {
datasketches::resize_factor to_cpp_tuple_resize_factor(TupleResizeFactor rf) {
switch (rf) {
case TupleResizeFactor::X1: return datasketches::resize_factor::X1;
case TupleResizeFactor::X2: return datasketches::resize_factor::X2;
case TupleResizeFactor::X4: return datasketches::resize_factor::X4;
case TupleResizeFactor::X8: return datasketches::resize_factor::X8;
default: throw std::invalid_argument("unknown TupleResizeFactor");
}
}
namespace {
datasketches::update_array_of_doubles_sketch build_sketch(uint8_t lg_k, TupleResizeFactor rf, float p, uint8_t num_values) {
datasketches::update_array_of_doubles_sketch::builder builder{
datasketches::default_array_of_doubles_update_policy(num_values)};
builder.set_lg_k(lg_k);
builder.set_resize_factor(to_cpp_tuple_resize_factor(rf));
builder.set_p(p);
return builder.build();
}
void check_values_len(const datasketches::update_array_of_doubles_sketch& sketch, rust::Slice<const double> values) {
if (values.size() != sketch.get_num_values()) {
throw std::invalid_argument("values length does not match num_values");
}
}
}
ArrayOfDoublesSketchShim::ArrayOfDoublesSketchShim(uint8_t lg_k, TupleResizeFactor rf, float p, uint8_t num_values)
: sketch_(build_sketch(lg_k, rf, p, num_values)) {}
void ArrayOfDoublesSketchShim::update_u64(uint64_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_i64(int64_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_u32(uint32_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_i32(int32_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_u16(uint16_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_i16(int16_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_u8(uint8_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_i8(int8_t key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_f64(double key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key, values.data());
}
void ArrayOfDoublesSketchShim::update_str(rust::Str key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(std::string(key), values.data());
}
void ArrayOfDoublesSketchShim::update_bytes(rust::Slice<const uint8_t> key, rust::Slice<const double> values) {
check_values_len(sketch_, values);
sketch_.update(key.data(), key.size(), values.data());
}
void ArrayOfDoublesSketchShim::trim() { sketch_.trim(); }
void ArrayOfDoublesSketchShim::reset() { sketch_.reset(); }
double ArrayOfDoublesSketchShim::get_estimate() const { return sketch_.get_estimate(); }
double ArrayOfDoublesSketchShim::get_lower_bound(uint8_t num_std_dev) const {
return sketch_.get_lower_bound(num_std_dev);
}
double ArrayOfDoublesSketchShim::get_upper_bound(uint8_t num_std_dev) const {
return sketch_.get_upper_bound(num_std_dev);
}
bool ArrayOfDoublesSketchShim::is_empty() const { return sketch_.is_empty(); }
bool ArrayOfDoublesSketchShim::is_estimation_mode() const { return sketch_.is_estimation_mode(); }
bool ArrayOfDoublesSketchShim::is_ordered() const { return sketch_.is_ordered(); }
double ArrayOfDoublesSketchShim::get_theta() const { return sketch_.get_theta(); }
uint32_t ArrayOfDoublesSketchShim::get_num_retained() const { return sketch_.get_num_retained(); }
uint8_t ArrayOfDoublesSketchShim::get_num_values() const { return sketch_.get_num_values(); }
rust::Vec<uint64_t> ArrayOfDoublesSketchShim::entry_hashes() const {
rust::Vec<uint64_t> out;
for (const auto& entry : sketch_) out.push_back(entry.first);
return out;
}
rust::Vec<double> ArrayOfDoublesSketchShim::entry_values() const {
rust::Vec<double> out;
for (const auto& entry : sketch_) {
for (uint8_t i = 0; i < entry.second.size(); ++i) out.push_back(entry.second[i]);
}
return out;
}
std::unique_ptr<ArrayOfDoublesSketchShim> new_array_of_doubles_sketch(uint8_t lg_k, TupleResizeFactor rf, float p, uint8_t num_values) {
return std::make_unique<ArrayOfDoublesSketchShim>(lg_k, rf, p, num_values);
}
std::unique_ptr<CompactArrayOfDoublesSketchShim> ArrayOfDoublesSketchShim::compact(bool ordered) const {
return array_of_doubles_sketch_compact(*this, ordered);
}
}