m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Suffix Tree
(string/suffix_tree.hpp)

Overview

SuffixTree stores all suffixes of a static text in a compact trie. It supports substring lookup, occurrence counting, representative occurrences, and direct tree traversal.

Construction uses Ukkonen’s algorithm. A unique terminal symbol is appended internally, so every suffix, including the empty suffix, ends at its own leaf. The terminal symbol has index AlphabetSize and cannot occur in input.

For a fixed alphabet, construction takes O(N) time, creates at most max(2, 2N + 1) nodes, and uses O(N * AlphabetSize) memory for fixed transition arrays.

How to Use the Tree

Node zero is the root. Every non-root node has one incoming edge, and node(v).parent is the node at the other end of that edge. The edge label is the slice

text.substr(tree.node(v).left, tree.edge_length(v))

except that a leaf edge may include the internal terminal symbol at position text.size(). In that case node(v).right == text.size() + 1; exclude the last position when reading the label from the original string.

There are two ways to access children:

// Follow a known character in O(1).
int child = tree.child(node, 'a');

// Visit only existing children in O(number of children).
tree.for_each_child(node, [&](int symbol, int child) {
    // symbol is 0 for 'a', 1 for 'b', ..., or terminal_symbol.
});

for_each_child does not scan the alphabet. Each node stores a linked list of its actual children in ascending symbol order. The same list can be traversed manually when a callback is inconvenient:

for (
    int child = tree.node(node).first_child;
    child != tree.null_node;
    child = tree.node(child).next_sibling
) {
    int symbol = tree.node(child).incoming_symbol;
    // Use child and symbol.
}

Concatenating incoming edge labels from the root to a node gives the string represented by that node. A leaf represents the suffix beginning at node(leaf).suffix_start. node(v).leaf_count is the number of occurrences of the path string of v.

Most substring-query code does not need to traverse the topology manually:

auto locus = tree.find(pattern);
if (locus) {
    int occurrences = tree.node(locus.node).leaf_count;
}

The pattern can end in the middle of an edge. locus.node is the child at the end of that edge, and locus.offset is the number of edge symbols consumed. The descendant leaves—and therefore the occurrence count—are the same for every point inside that edge.

Template Parameters

Every input symbol c must satisfy FirstCharacter <= c < FirstCharacter + AlphabetSize. For decimal strings, use SuffixTree<10, '0'>.

Node and Locus Fields

An edge into node v is labeled by the internal augmented-text interval [node(v).left, node(v).right). The augmented text has length N + 1; position N is the unique terminal symbol.

Field Meaning
next[c] Child whose edge starts with symbol index c, or null_node. Index AlphabetSize is the terminal symbol.
suffix_link Ukkonen suffix link for an internal node. It may be null_node for a leaf.
parent Parent node, or null_node at the root.
left, right Half-open augmented-text interval labeling the incoming edge.
suffix_start Starting position of the represented suffix for a leaf, or -1 for an internal node.
representative_suffix Start of one descendant suffix.
leaf_count Number of descendant leaves.
incoming_symbol First symbol index of the incoming edge, or -1 at the root.
first_child First actual child in symbol order, or null_node.
next_sibling Next child of the same parent, or null_node.
child_count Number of actual children.

Locus contains node and offset. A successful search may end inside the incoming edge of node; offset is the number of consumed symbols on that edge. It is an explicit node exactly when offset == edge_length(node). A Locus converts to false when no match exists.

Methods

Let V be the number of nodes, L a query length, and A = AlphabetSize.

Method Description Complexity
SuffixTree() Builds the tree of the empty text and its terminal suffix. O(A)
template<class Sequence> explicit SuffixTree(const Sequence& sequence) Builds the suffix tree of sequence. O(N * A)
int size() const, int node_count() const Returns V, including the root and terminal leaf. O(1)
bool empty() const Returns whether the original text is empty. O(1)
int text_length() const Returns N, excluding the terminal symbol. O(1)
node_id root() const Returns node zero. O(1)
const Node& node(node_id id) const Returns node metadata. O(1)
const std::vector<Node>& nodes() const Returns all nodes. O(1)
int edge_length(node_id id) const Returns the length of the incoming edge. O(1)
bool is_leaf(node_id id) const Tests whether a node represents one complete suffix. O(1)
template<class Symbol> node_id child(node_id id, const Symbol& symbol) const Returns an input-symbol child, or null_node. O(1)
node_id child_by_index(node_id id, int symbol) const Returns a child by symbol index, including terminal_symbol. O(1)
template<class Callback> void for_each_child(node_id id, Callback callback) const Calls callback(symbol, child) for each actual child in symbol-index order. O(node(id).child_count)
void clear() Replaces the tree with the empty-text tree. O(V + A)
template<class Sequence> void build(const Sequence& sequence) Replaces the tree with the suffix tree of sequence. O(V + N * A)
template<class Sequence> Locus find(const Sequence& sequence) const Returns the locus of a substring, or a false locus. O(L)
template<class Sequence> bool contains(const Sequence& sequence) const Tests whether a sequence is a substring. O(L)
template<class Sequence> int count_occurrences(const Sequence& sequence) const Counts possibly overlapping occurrences. O(L)
template<class Sequence> std::pair<int, int> representative_occurrence(const Sequence& sequence) const Returns one half-open occurrence, or (-1, -1). O(L)
long long distinct_substring_count() const Counts distinct nonempty substrings. O(V)

For the empty query, count_occurrences returns N + 1 and representative_occurrence returns an empty interval.

Node handles remain valid until build or clear. Both operations rebuild the whole tree and invalidate all earlier handles and references.

Example

#include "string/suffix_tree.hpp"
#include <algorithm>
#include <iostream>
#include <string>
#include <vector>

int main() {
    std::string text = "ababa";
    m1une::string::SuffixTree<> tree(text);

    std::cout << tree.contains(std::string("bab")) << '\n';       // 1
    std::cout << tree.count_occurrences(std::string("aba")) << '\n';  // 2
    std::cout << tree.distinct_substring_count() << '\n';         // 9

    auto occurrence = tree.representative_occurrence(std::string("bab"));
    std::cout << occurrence.first << ' ' << occurrence.second << '\n';

    // Print every edge as: parent, child, label.
    std::vector<int> stack(1, tree.root());
    while (!stack.empty()) {
        int parent = stack.back();
        stack.pop_back();
        tree.for_each_child(parent, [&](int symbol, int child) {
            const auto& current = tree.node(child);
            int right = std::min(current.right, tree.text_length());
            std::string label = text.substr(current.left, right - current.left);
            std::cout << parent << " -> " << child << ": " << label;
            if (symbol == tree.terminal_symbol || current.right > tree.text_length()) {
                std::cout << '$';
            }
            std::cout << '\n';
            stack.push_back(child);
        });
    }
}

Required by

Verified with

Code

#ifndef M1UNE_STRING_SUFFIX_TREE_HPP
#define M1UNE_STRING_SUFFIX_TREE_HPP 1

#include <algorithm>
#include <array>
#include <cassert>
#include <cstddef>
#include <limits>
#include <utility>
#include <vector>

namespace m1une {
namespace string {

template <int AlphabetSize = 26, int FirstCharacter = 'a'>
struct SuffixTree {
    static_assert(0 < AlphabetSize);

    using node_id = int;
    static constexpr node_id root_node = 0;
    static constexpr node_id null_node = -1;
    static constexpr int terminal_symbol = AlphabetSize;

    struct Node {
        std::array<node_id, AlphabetSize + 1> next;
        node_id suffix_link;
        node_id parent;
        int left;
        int right;
        int suffix_start;
        int representative_suffix;
        int leaf_count;
        int incoming_symbol;
        node_id first_child;
        node_id next_sibling;
        int child_count;

        Node(int left_value = 0, int right_value = 0, node_id parent_value = null_node)
            : suffix_link(null_node),
              parent(parent_value),
              left(left_value),
              right(right_value),
              suffix_start(-1),
              representative_suffix(-1),
              leaf_count(0),
              incoming_symbol(-1),
              first_child(null_node),
              next_sibling(null_node),
              child_count(0) {
            next.fill(null_node);
        }
    };

    struct Locus {
        node_id node;
        int offset;

        explicit operator bool() const {
            return node != null_node;
        }

        friend bool operator==(const Locus&, const Locus&) = default;
    };

   private:
    struct ActivePoint {
        node_id node;
        int offset;
    };

    std::vector<Node> _nodes;
    std::vector<int> _text;
    ActivePoint _active;
    int _text_length;

    template <class Symbol>
    static int symbol_index(const Symbol& symbol) {
        int index = int(symbol) - FirstCharacter;
        assert(0 <= index && index < AlphabetSize);
        return index;
    }

    int edge_length_unchecked(node_id id) const {
        return _nodes[id].right - _nodes[id].left;
    }

    node_id new_node(int left, int right, node_id parent) {
        assert(_nodes.size() < std::size_t(std::numeric_limits<int>::max()));
        _nodes.emplace_back(left, right, parent);
        return int(_nodes.size()) - 1;
    }

    ActivePoint go(ActivePoint point, int left, int right) const {
        while (left < right) {
            if (point.offset == edge_length_unchecked(point.node)) {
                point = {_nodes[point.node].next[_text[left]], 0};
                if (point.node == null_node) return point;
            } else {
                if (_text[_nodes[point.node].left + point.offset] != _text[left]) {
                    return {null_node, 0};
                }
                int remaining = edge_length_unchecked(point.node) - point.offset;
                if (right - left < remaining) {
                    point.offset += right - left;
                    return point;
                }
                left += remaining;
                point.offset = edge_length_unchecked(point.node);
            }
        }
        return point;
    }

    node_id split(ActivePoint point) {
        if (point.offset == edge_length_unchecked(point.node)) return point.node;
        if (point.offset == 0) return _nodes[point.node].parent;

        node_id child = point.node;
        node_id parent = _nodes[child].parent;
        int left = _nodes[child].left;
        node_id middle = new_node(left, left + point.offset, parent);
        _nodes[parent].next[_text[left]] = middle;
        _nodes[middle].next[_text[left + point.offset]] = child;
        _nodes[child].parent = middle;
        _nodes[child].left += point.offset;
        return middle;
    }

    node_id get_suffix_link(node_id id) {
        if (_nodes[id].suffix_link != null_node) return _nodes[id].suffix_link;
        node_id parent = _nodes[id].parent;
        if (parent == null_node) return root_node;

        node_id parent_link = get_suffix_link(parent);
        ActivePoint point = {
            parent_link,
            edge_length_unchecked(parent_link)
        };
        int left = _nodes[id].left + (parent == root_node);
        point = go(point, left, _nodes[id].right);
        assert(point.node != null_node);
        return _nodes[id].suffix_link = split(point);
    }

    void extend(int position) {
        while (true) {
            ActivePoint next = go(_active, position, position + 1);
            if (next.node != null_node) {
                _active = next;
                return;
            }

            node_id middle = split(_active);
            node_id leaf = new_node(position, int(_text.size()), middle);
            _nodes[middle].next[_text[position]] = leaf;

            _active.node = get_suffix_link(middle);
            _active.offset = edge_length_unchecked(_active.node);
            if (middle == root_node) return;
        }
    }

    void finish_metadata() {
        std::vector<node_id> order;
        order.reserve(_nodes.size());
        order.push_back(root_node);
        std::vector<int> depth(_nodes.size(), 0);

        for (std::size_t i = 0; i < order.size(); i++) {
            node_id id = order[i];
            node_id previous_child = null_node;
            for (int symbol = 0; symbol <= terminal_symbol; symbol++) {
                node_id child = _nodes[id].next[symbol];
                if (child == null_node) continue;
                _nodes[child].incoming_symbol = symbol;
                if (previous_child == null_node) {
                    _nodes[id].first_child = child;
                } else {
                    _nodes[previous_child].next_sibling = child;
                }
                previous_child = child;
                _nodes[id].child_count++;
                depth[child] = depth[id] + edge_length_unchecked(child);
                order.push_back(child);
            }
        }

        for (int i = int(order.size()) - 1; i >= 0; i--) {
            node_id id = order[i];
            bool leaf = true;
            for (node_id child : _nodes[id].next) {
                if (child == null_node) continue;
                leaf = false;
                _nodes[id].leaf_count += _nodes[child].leaf_count;
                if (_nodes[id].representative_suffix == -1) {
                    _nodes[id].representative_suffix = _nodes[child].representative_suffix;
                }
            }
            if (leaf) {
                _nodes[id].suffix_start = int(_text.size()) - depth[id];
                _nodes[id].representative_suffix = _nodes[id].suffix_start;
                _nodes[id].leaf_count = 1;
            }
        }
    }

    void initialize() {
        _nodes.clear();
        _nodes.reserve(2 * _text.size() + 1);
        _nodes.emplace_back();
        _nodes[root_node].suffix_link = root_node;
        _active = {root_node, 0};
        for (int position = 0; position < int(_text.size()); position++) extend(position);
        finish_metadata();
    }

   public:
    SuffixTree() {
        clear();
    }

    template <class Sequence>
    explicit SuffixTree(const Sequence& sequence) {
        build(sequence);
    }

    int size() const {
        return node_count();
    }

    bool empty() const {
        return _text_length == 0;
    }

    int node_count() const {
        return int(_nodes.size());
    }

    int text_length() const {
        return _text_length;
    }

    node_id root() const {
        return root_node;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && id < node_count());
        return _nodes[id];
    }

    const std::vector<Node>& nodes() const {
        return _nodes;
    }

    int edge_length(node_id id) const {
        assert(0 <= id && id < node_count());
        return edge_length_unchecked(id);
    }

    bool is_leaf(node_id id) const {
        assert(0 <= id && id < node_count());
        return _nodes[id].suffix_start != -1;
    }

    template <class Symbol>
    node_id child(node_id id, const Symbol& symbol) const {
        assert(0 <= id && id < node_count());
        return _nodes[id].next[symbol_index(symbol)];
    }

    node_id child_by_index(node_id id, int symbol) const {
        assert(0 <= id && id < node_count());
        assert(0 <= symbol && symbol <= terminal_symbol);
        return _nodes[id].next[symbol];
    }

    template <class Callback>
    void for_each_child(node_id id, Callback callback) const {
        assert(0 <= id && id < node_count());
        for (
            node_id child_id = _nodes[id].first_child;
            child_id != null_node;
            child_id = _nodes[child_id].next_sibling
        ) {
            callback(_nodes[child_id].incoming_symbol, child_id);
        }
    }

    void clear() {
        _text.clear();
        _text.push_back(terminal_symbol);
        _text_length = 0;
        initialize();
    }

    template <class Sequence>
    void build(const Sequence& sequence) {
        _text.clear();
        for (const auto& symbol : sequence) _text.push_back(symbol_index(symbol));
        assert(_text.size() < std::size_t(std::numeric_limits<int>::max()));
        _text_length = int(_text.size());
        _text.push_back(terminal_symbol);
        initialize();
    }

    template <class Sequence>
    Locus find(const Sequence& sequence) const {
        ActivePoint point = {root_node, 0};
        for (const auto& value : sequence) {
            int symbol = symbol_index(value);
            if (point.offset == edge_length_unchecked(point.node)) {
                point = {_nodes[point.node].next[symbol], 0};
                if (point.node == null_node) return {null_node, 0};
            }
            if (_text[_nodes[point.node].left + point.offset] != symbol) {
                return {null_node, 0};
            }
            point.offset++;
        }
        return {point.node, point.offset};
    }

    template <class Sequence>
    bool contains(const Sequence& sequence) const {
        return bool(find(sequence));
    }

    template <class Sequence>
    int count_occurrences(const Sequence& sequence) const {
        Locus locus = find(sequence);
        return locus ? _nodes[locus.node].leaf_count : 0;
    }

    template <class Sequence>
    std::pair<int, int> representative_occurrence(const Sequence& sequence) const {
        Locus locus = {root_node, 0};
        int length = 0;
        for (const auto& value : sequence) {
            int symbol = symbol_index(value);
            if (locus.offset == edge_length_unchecked(locus.node)) {
                locus = {_nodes[locus.node].next[symbol], 0};
                if (locus.node == null_node) return {-1, -1};
            }
            if (_text[_nodes[locus.node].left + locus.offset] != symbol) return {-1, -1};
            locus.offset++;
            length++;
        }
        int left = _nodes[locus.node].representative_suffix;
        return {left, left + length};
    }

    long long distinct_substring_count() const {
        long long result = 0;
        for (node_id id = 1; id < node_count(); id++) {
            result += std::max(0, std::min(_nodes[id].right, _text_length) - _nodes[id].left);
        }
        return result;
    }
};

}  // namespace string
}  // namespace m1une

#endif  // M1UNE_STRING_SUFFIX_TREE_HPP
#line 1 "string/suffix_tree.hpp"



#include <algorithm>
#include <array>
#include <cassert>
#include <cstddef>
#include <limits>
#include <utility>
#include <vector>

namespace m1une {
namespace string {

template <int AlphabetSize = 26, int FirstCharacter = 'a'>
struct SuffixTree {
    static_assert(0 < AlphabetSize);

    using node_id = int;
    static constexpr node_id root_node = 0;
    static constexpr node_id null_node = -1;
    static constexpr int terminal_symbol = AlphabetSize;

    struct Node {
        std::array<node_id, AlphabetSize + 1> next;
        node_id suffix_link;
        node_id parent;
        int left;
        int right;
        int suffix_start;
        int representative_suffix;
        int leaf_count;
        int incoming_symbol;
        node_id first_child;
        node_id next_sibling;
        int child_count;

        Node(int left_value = 0, int right_value = 0, node_id parent_value = null_node)
            : suffix_link(null_node),
              parent(parent_value),
              left(left_value),
              right(right_value),
              suffix_start(-1),
              representative_suffix(-1),
              leaf_count(0),
              incoming_symbol(-1),
              first_child(null_node),
              next_sibling(null_node),
              child_count(0) {
            next.fill(null_node);
        }
    };

    struct Locus {
        node_id node;
        int offset;

        explicit operator bool() const {
            return node != null_node;
        }

        friend bool operator==(const Locus&, const Locus&) = default;
    };

   private:
    struct ActivePoint {
        node_id node;
        int offset;
    };

    std::vector<Node> _nodes;
    std::vector<int> _text;
    ActivePoint _active;
    int _text_length;

    template <class Symbol>
    static int symbol_index(const Symbol& symbol) {
        int index = int(symbol) - FirstCharacter;
        assert(0 <= index && index < AlphabetSize);
        return index;
    }

    int edge_length_unchecked(node_id id) const {
        return _nodes[id].right - _nodes[id].left;
    }

    node_id new_node(int left, int right, node_id parent) {
        assert(_nodes.size() < std::size_t(std::numeric_limits<int>::max()));
        _nodes.emplace_back(left, right, parent);
        return int(_nodes.size()) - 1;
    }

    ActivePoint go(ActivePoint point, int left, int right) const {
        while (left < right) {
            if (point.offset == edge_length_unchecked(point.node)) {
                point = {_nodes[point.node].next[_text[left]], 0};
                if (point.node == null_node) return point;
            } else {
                if (_text[_nodes[point.node].left + point.offset] != _text[left]) {
                    return {null_node, 0};
                }
                int remaining = edge_length_unchecked(point.node) - point.offset;
                if (right - left < remaining) {
                    point.offset += right - left;
                    return point;
                }
                left += remaining;
                point.offset = edge_length_unchecked(point.node);
            }
        }
        return point;
    }

    node_id split(ActivePoint point) {
        if (point.offset == edge_length_unchecked(point.node)) return point.node;
        if (point.offset == 0) return _nodes[point.node].parent;

        node_id child = point.node;
        node_id parent = _nodes[child].parent;
        int left = _nodes[child].left;
        node_id middle = new_node(left, left + point.offset, parent);
        _nodes[parent].next[_text[left]] = middle;
        _nodes[middle].next[_text[left + point.offset]] = child;
        _nodes[child].parent = middle;
        _nodes[child].left += point.offset;
        return middle;
    }

    node_id get_suffix_link(node_id id) {
        if (_nodes[id].suffix_link != null_node) return _nodes[id].suffix_link;
        node_id parent = _nodes[id].parent;
        if (parent == null_node) return root_node;

        node_id parent_link = get_suffix_link(parent);
        ActivePoint point = {
            parent_link,
            edge_length_unchecked(parent_link)
        };
        int left = _nodes[id].left + (parent == root_node);
        point = go(point, left, _nodes[id].right);
        assert(point.node != null_node);
        return _nodes[id].suffix_link = split(point);
    }

    void extend(int position) {
        while (true) {
            ActivePoint next = go(_active, position, position + 1);
            if (next.node != null_node) {
                _active = next;
                return;
            }

            node_id middle = split(_active);
            node_id leaf = new_node(position, int(_text.size()), middle);
            _nodes[middle].next[_text[position]] = leaf;

            _active.node = get_suffix_link(middle);
            _active.offset = edge_length_unchecked(_active.node);
            if (middle == root_node) return;
        }
    }

    void finish_metadata() {
        std::vector<node_id> order;
        order.reserve(_nodes.size());
        order.push_back(root_node);
        std::vector<int> depth(_nodes.size(), 0);

        for (std::size_t i = 0; i < order.size(); i++) {
            node_id id = order[i];
            node_id previous_child = null_node;
            for (int symbol = 0; symbol <= terminal_symbol; symbol++) {
                node_id child = _nodes[id].next[symbol];
                if (child == null_node) continue;
                _nodes[child].incoming_symbol = symbol;
                if (previous_child == null_node) {
                    _nodes[id].first_child = child;
                } else {
                    _nodes[previous_child].next_sibling = child;
                }
                previous_child = child;
                _nodes[id].child_count++;
                depth[child] = depth[id] + edge_length_unchecked(child);
                order.push_back(child);
            }
        }

        for (int i = int(order.size()) - 1; i >= 0; i--) {
            node_id id = order[i];
            bool leaf = true;
            for (node_id child : _nodes[id].next) {
                if (child == null_node) continue;
                leaf = false;
                _nodes[id].leaf_count += _nodes[child].leaf_count;
                if (_nodes[id].representative_suffix == -1) {
                    _nodes[id].representative_suffix = _nodes[child].representative_suffix;
                }
            }
            if (leaf) {
                _nodes[id].suffix_start = int(_text.size()) - depth[id];
                _nodes[id].representative_suffix = _nodes[id].suffix_start;
                _nodes[id].leaf_count = 1;
            }
        }
    }

    void initialize() {
        _nodes.clear();
        _nodes.reserve(2 * _text.size() + 1);
        _nodes.emplace_back();
        _nodes[root_node].suffix_link = root_node;
        _active = {root_node, 0};
        for (int position = 0; position < int(_text.size()); position++) extend(position);
        finish_metadata();
    }

   public:
    SuffixTree() {
        clear();
    }

    template <class Sequence>
    explicit SuffixTree(const Sequence& sequence) {
        build(sequence);
    }

    int size() const {
        return node_count();
    }

    bool empty() const {
        return _text_length == 0;
    }

    int node_count() const {
        return int(_nodes.size());
    }

    int text_length() const {
        return _text_length;
    }

    node_id root() const {
        return root_node;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && id < node_count());
        return _nodes[id];
    }

    const std::vector<Node>& nodes() const {
        return _nodes;
    }

    int edge_length(node_id id) const {
        assert(0 <= id && id < node_count());
        return edge_length_unchecked(id);
    }

    bool is_leaf(node_id id) const {
        assert(0 <= id && id < node_count());
        return _nodes[id].suffix_start != -1;
    }

    template <class Symbol>
    node_id child(node_id id, const Symbol& symbol) const {
        assert(0 <= id && id < node_count());
        return _nodes[id].next[symbol_index(symbol)];
    }

    node_id child_by_index(node_id id, int symbol) const {
        assert(0 <= id && id < node_count());
        assert(0 <= symbol && symbol <= terminal_symbol);
        return _nodes[id].next[symbol];
    }

    template <class Callback>
    void for_each_child(node_id id, Callback callback) const {
        assert(0 <= id && id < node_count());
        for (
            node_id child_id = _nodes[id].first_child;
            child_id != null_node;
            child_id = _nodes[child_id].next_sibling
        ) {
            callback(_nodes[child_id].incoming_symbol, child_id);
        }
    }

    void clear() {
        _text.clear();
        _text.push_back(terminal_symbol);
        _text_length = 0;
        initialize();
    }

    template <class Sequence>
    void build(const Sequence& sequence) {
        _text.clear();
        for (const auto& symbol : sequence) _text.push_back(symbol_index(symbol));
        assert(_text.size() < std::size_t(std::numeric_limits<int>::max()));
        _text_length = int(_text.size());
        _text.push_back(terminal_symbol);
        initialize();
    }

    template <class Sequence>
    Locus find(const Sequence& sequence) const {
        ActivePoint point = {root_node, 0};
        for (const auto& value : sequence) {
            int symbol = symbol_index(value);
            if (point.offset == edge_length_unchecked(point.node)) {
                point = {_nodes[point.node].next[symbol], 0};
                if (point.node == null_node) return {null_node, 0};
            }
            if (_text[_nodes[point.node].left + point.offset] != symbol) {
                return {null_node, 0};
            }
            point.offset++;
        }
        return {point.node, point.offset};
    }

    template <class Sequence>
    bool contains(const Sequence& sequence) const {
        return bool(find(sequence));
    }

    template <class Sequence>
    int count_occurrences(const Sequence& sequence) const {
        Locus locus = find(sequence);
        return locus ? _nodes[locus.node].leaf_count : 0;
    }

    template <class Sequence>
    std::pair<int, int> representative_occurrence(const Sequence& sequence) const {
        Locus locus = {root_node, 0};
        int length = 0;
        for (const auto& value : sequence) {
            int symbol = symbol_index(value);
            if (locus.offset == edge_length_unchecked(locus.node)) {
                locus = {_nodes[locus.node].next[symbol], 0};
                if (locus.node == null_node) return {-1, -1};
            }
            if (_text[_nodes[locus.node].left + locus.offset] != symbol) return {-1, -1};
            locus.offset++;
            length++;
        }
        int left = _nodes[locus.node].representative_suffix;
        return {left, left + length};
    }

    long long distinct_substring_count() const {
        long long result = 0;
        for (node_id id = 1; id < node_count(); id++) {
            result += std::max(0, std::min(_nodes[id].right, _text_length) - _nodes[id].left);
        }
        return result;
    }
};

}  // namespace string
}  // namespace m1une
Back to top page