#pragma once
#include <cstdint>
#include <map>
#include <set>
#include <string>
#include <string_view>
#include <vector>
struct common_trie {
struct node {
std::map<uint32_t, size_t> children;
int32_t pattern = -1;
};
std::vector<node> nodes;
common_trie() {
create_node();
}
common_trie(const std::vector<std::string> & words) : common_trie() {
for (const auto & w : words) {
insert(w);
}
}
enum match_result { NO_MATCH, PARTIAL_MATCH, COMPLETE_MATCH };
match_result check_at(std::string_view sv, size_t start_pos) const;
int32_t insert(const std::string & word);
int32_t insert(const std::vector<uint32_t> & symbols);
private:
int32_t n_patterns = 0;
size_t create_node() {
size_t index = nodes.size();
nodes.emplace_back();
return index;
}
};
struct common_aho_corasick {
common_trie t;
std::vector<size_t> fail;
std::vector<size_t> order;
std::vector<int32_t> match;
std::set<uint32_t> alphabet;
common_aho_corasick(common_trie trie);
common_aho_corasick(const std::vector<std::string> & strings)
: common_aho_corasick(common_trie(strings)) {}
size_t num_states() const { return t.nodes.size(); }
bool is_terminal(size_t s) const { return match[s] >= 0; }
int32_t match_pattern(size_t s) const { return match[s]; }
size_t next(size_t state, uint32_t ch) const;
};