#include <cmath>
#include <cstdint>
#include <iostream>
#include <stdexcept>
#include <string>
#include <utility>
#include <vector>
namespace util {
class BloomFilter {
public:
BloomFilter(std::uint64_t expected_n, double fpr = 0.01) {
if (expected_n == 0 || !(fpr > 0.0 && fpr < 1.0)) {
throw std::invalid_argument("invalid BloomFilter params");
}
const double ln2 = std::log(2.0);
const std::uint64_t bits = static_cast<std::uint64_t>(
-static_cast<double>(expected_n) * std::log(fpr) / (ln2 * ln2) + 0.5);
m_ = next_pow2(bits < 64 ? 64 : bits);
mask_ = m_ - 1;
std::uint64_t hashes = static_cast<std::uint64_t>(
static_cast<double>(m_) / expected_n * ln2 + 0.5);
if (hashes < 1) hashes = 1;
if (hashes > 16) hashes = 16;
k_ = hashes;
words_.assign(m_ >> 6, 0);
}
void add(std::uint64_t key) { insert(hash64(key)); }
void add(const std::string& key) { insert(hash64(key)); }
bool may_contain(std::uint64_t key) const { return probe(hash64(key)); }
bool may_contain(const std::string& key) const { return probe(hash64(key)); }
std::uint64_t bit_count() const { return m_; }
std::uint64_t hash_count() const { return k_; }
std::size_t memory_bytes() const { return words_.size() * 8; }
private:
std::uint64_t m_;
std::uint64_t mask_;
std::uint64_t k_;
std::vector<std::uint64_t> words_;
static std::uint64_t next_pow2(std::uint64_t x) {
--x;
x |= x >> 1;
x |= x >> 2;
x |= x >> 4;
x |= x >> 8;
x |= x >> 16;
x |= x >> 32;
return x + 1;
}
void set_bit(std::uint64_t i) { words_[i >> 6] |= 1ULL << (i & 63); }
bool test_bit(std::uint64_t i) const {
return (words_[i >> 6] >> (i & 63)) & 1ULL;
}
std::uint64_t slot(std::uint64_t h1, std::uint64_t h2, std::uint64_t i) const {
return (h1 + i * (h2 | 1ULL)) & mask_;
}
void insert(std::pair<std::uint64_t, std::uint64_t> h) {
for (std::uint64_t i = 0; i < k_; ++i) {
set_bit(slot(h.first, h.second, i));
}
}
bool probe(std::pair<std::uint64_t, std::uint64_t> h) const {
for (std::uint64_t i = 0; i < k_; ++i) {
if (!test_bit(slot(h.first, h.second, i))) return false;
}
return true;
}
static std::uint64_t mix64(std::uint64_t x) {
x += 0x9e3779b97f4a7c15ULL;
x = (x ^ (x >> 30)) * 0xbf58476d1ce4e5b9ULL;
x = (x ^ (x >> 27)) * 0x94d049bb133111ebULL;
return x ^ (x >> 31);
}
static std::pair<std::uint64_t, std::uint64_t> hash64(std::uint64_t key) {
return std::make_pair(mix64(key), mix64(key ^ 0x9e3779b97f4a7c15ULL));
}
static std::pair<std::uint64_t, std::uint64_t> hash64(const std::string& s) {
std::uint64_t h1 = 0xcbf29ce484222325ULL;
std::uint64_t h2 = 0x9e3779b97f4a7c15ULL;
for (std::size_t i = 0; i < s.size(); ++i) {
unsigned char c = static_cast<unsigned char>(s[i]);
h1 = mix64(h1 ^ c);
h2 = mix64(h2 ^ c);
}
return std::make_pair(mix64(h1 ^ s.size()), mix64(h2 ^ s.size()));
}
};
} // namespace util
int main() {
const std::uint64_t N = 100000;
util::BloomFilter bf(N, 0.01);
std::cout << "m=" << bf.bit_count()
<< " k=" << bf.hash_count()
<< " bytes=" << bf.memory_bytes() << "\n";
for (std::uint64_t i = 0; i < N; ++i) bf.add(i);
std::uint64_t fp = 0;
for (std::uint64_t i = N; i < 2 * N; ++i) {
if (bf.may_contain(i)) ++fp;
}
std::cout << "has 42=" << bf.may_contain(static_cast<std::uint64_t>(42)) << "\n";
std::cout << "has 999999=" << bf.may_contain(static_cast<std::uint64_t>(999999)) << "\n";
std::cout << "false_positives=" << fp << "/" << N << "\n";
}