m1une's library

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

View on GitHub

:heavy_check_mark: Rollback Ordered Set
(ds/bst/rollback_ordered_set.hpp)

Overview

RollbackOrderedSet<T, Compare> is an ordered set with order statistics and registered-snapshot rollback. It stores its mutable ordered tree independently of any versioned structure. Compare must define a strict weak ordering, as for std::set.

Methods

Constructors and read-only search, order-statistic, conversion, and split methods follow the corresponding mutable structure.

Method Description Complexity
void clear() Clears the set. $O(N)$
bool insert(T key) Inserts a key, reports whether it was new. $O(\log N)$
bool erase(const T& key) Erases a key, reports whether it existed. $O(\log N)$
void merge(const RollbackOrderedSet& other) Inserts an ordered, disjoint set of size $M$. $O(M \log(N+M))$
int snapshot() Registers the current state and returns its token. $O(1)$
int snapshot_count() const Returns the number of active snapshots. $O(1)$
void reserve_snapshots(int count) Reserves snapshot tokens. $O(H)$
void rollback(int state) Rolls back to a current-path snapshot. $O(F)$ total
void clear_history(), void release() Releases saved states, or all states. $O(F)$

Snapshot semantics

Updates made before the first snapshot() retain no rollback data. A snapshot token is positive and valid only on the current path. rollback(state) restores that registered state, keeps it active, and invalidates newer snapshots. clear_history() commits the current state and invalidates every token. No per-update reversal operation is provided.

Example

#include "ds/bst/rollback_ordered_set.hpp"

m1une::ds::RollbackOrderedSet<int> set;
set.insert(4);
int state = set.snapshot();
set.insert(2);
set.rollback(state);
assert(!set.contains(2));

Depends on

Verified with

Code

#ifndef M1UNE_DS_BST_ROLLBACK_ORDERED_SET_HPP
#define M1UNE_DS_BST_ROLLBACK_ORDERED_SET_HPP 1

#include <cassert>
#include <functional>
#include <initializer_list>
#include <optional>
#include <utility>
#include <vector>

#include "ordered_set.hpp"

namespace m1une {
namespace ds {

template <class T, class Compare = std::less<T>>
struct RollbackOrderedSet {
   private:
    enum class Kind { insert, erase, clear, merge };
    struct Entry {
        Kind kind;
        bool changed;
        std::optional<T> key;
        std::vector<T> keys;
    };

    OrderedSet<T, Compare> _data;
    std::vector<Entry> _history;
    std::vector<std::size_t> _checkpoints;

   public:
    explicit RollbackOrderedSet(Compare compare)
        : _data(std::move(compare)) {}
    RollbackOrderedSet() = default;

    RollbackOrderedSet(
        std::initializer_list<T> init,
        Compare compare = Compare()
    ) : _data(init, std::move(compare)) {}

    template <class Iterator>
    RollbackOrderedSet(
        Iterator first,
        Iterator last,
        Compare compare = Compare()
    ) : _data(first, last, std::move(compare)) {}

    int size() const { return _data.size(); }
    int unique_size() const { return _data.size(); }
    bool empty() const { return _data.empty(); }
    std::size_t node_count() const { return std::size_t(size()); }

    void clear() {
        if (_checkpoints.empty()) {
            _data.clear();
            return;
        }
        Entry entry{Kind::clear, !empty(), std::nullopt, {}};
        if (!empty()) entry.keys = _data.to_vector();
        _data.clear();
        _history.push_back(std::move(entry));
    }

    bool insert(T key) {
        if (_checkpoints.empty()) return _data.insert(std::move(key));
        bool changed = !_data.contains(key);
        Entry entry{Kind::insert, changed, std::nullopt, {}};
        if (changed) entry.key.emplace(key);
        _data.insert(std::move(key));
        _history.push_back(std::move(entry));
        return changed;
    }

    bool erase(const T& key) {
        if (_checkpoints.empty()) return _data.erase(key);
        bool changed = _data.contains(key);
        Entry entry{Kind::erase, changed, std::nullopt, {}};
        if (changed) entry.key.emplace(key);
        _data.erase(key);
        _history.push_back(std::move(entry));
        return changed;
    }

    void merge(const RollbackOrderedSet& other) {
        std::vector<T> keys = other.to_vector();
        for (const T& key : keys) {
            bool inserted = _data.insert(key);
            assert(inserted);
        }
        if (!_checkpoints.empty()) {
            _history.push_back(Entry{Kind::merge, !keys.empty(), std::nullopt, std::move(keys)});
        }
    }

    void merge(const OrderedSet<T, Compare>& other) {
        std::vector<T> keys = other.to_vector();
        for (const T& key : keys) {
            bool inserted = _data.insert(key);
            assert(inserted);
        }
        if (!_checkpoints.empty()) {
            _history.push_back(Entry{Kind::merge, !keys.empty(), std::nullopt, std::move(keys)});
        }
    }

    bool contains(const T& key) const { return _data.contains(key); }
    int count(const T& key) const { return _data.count(key); }
    const T* find_by_order(int order) const { return _data.find_by_order(order); }
    T kth(int order) const { return _data.kth(order); }
    int order_of_key(const T& key) const { return _data.order_of_key(key); }
    int count_less(const T& key) const { return _data.count_less(key); }
    int count_less_equal(const T& key) const { return _data.count_less_equal(key); }
    int count_greater(const T& key) const { return _data.count_greater(key); }
    int count_greater_equal(const T& key) const { return _data.count_greater_equal(key); }
    const T* lower_bound(const T& key) const { return _data.lower_bound(key); }
    const T* upper_bound(const T& key) const { return _data.upper_bound(key); }
    const T* min_ge(const T& key) const { return _data.min_ge(key); }
    const T* min_gt(const T& key) const { return _data.min_gt(key); }
    const T* max_le(const T& key) const { return _data.max_le(key); }
    const T* max_lt(const T& key) const { return _data.max_lt(key); }
    const T* min() const { return _data.min(); }
    const T* max() const { return _data.max(); }
    std::vector<T> to_vector() const { return _data.to_vector(); }

    int snapshot() { _checkpoints.push_back(_history.size()); return int(_checkpoints.size()); }
    int snapshot_count() const { return int(_checkpoints.size()); }

    void reserve_snapshots(int count) {
        assert(0 <= count);
        _checkpoints.reserve(count);
    }

   private:
    void restore_one() {
        Entry entry = std::move(_history.back());
        _history.pop_back();
        if (!entry.changed) return;
        if (entry.kind == Kind::insert) {
            _data.erase(*entry.key);
        } else if (entry.kind == Kind::erase) {
            _data.insert(std::move(*entry.key));
        } else if (entry.kind == Kind::clear) {
            for (T& key : entry.keys) _data.insert(std::move(key));
        } else {
            for (const T& key : entry.keys) _data.erase(key);
        }
    }

   public:

    void rollback(int state) {
        assert(1 <= state && state <= snapshot_count());
        while (_history.size() > _checkpoints[state - 1]) restore_one();
        _checkpoints.resize(state);
    }

    void clear_history() { _history.clear(); _checkpoints.clear(); }

    void release() {
        _data.clear();
        _history.clear();
        _checkpoints.clear();
    }
};

}  // namespace ds
}  // namespace m1une

#endif  // M1UNE_DS_BST_ROLLBACK_ORDERED_SET_HPP
#line 1 "ds/bst/rollback_ordered_set.hpp"



#include <cassert>
#include <functional>
#include <initializer_list>
#include <optional>
#include <utility>
#include <vector>

#line 1 "ds/bst/ordered_set.hpp"



#line 7 "ds/bst/ordered_set.hpp"
#include <memory>
#line 10 "ds/bst/ordered_set.hpp"

namespace m1une {
namespace ds {

template <typename T, typename Compare = std::less<T>>
struct OrderedSet {
   private:
    struct Node {
        T key;
        int size;
        Node* l;
        Node* r;

        explicit Node(T value)
            : key(std::move(value)), size(1), l(nullptr), r(nullptr) {}
    };

    static constexpr int pool_block_size = 1 << 15;

    struct Pool {
        std::vector<std::vector<Node>> blocks;
        std::vector<Node*> free_nodes;

        template <class... Args>
        Node* emplace(Args&&... args) {
            if (!free_nodes.empty()) {
                Node* result = free_nodes.back();
                free_nodes.pop_back();
                std::destroy_at(result);
                std::construct_at(result, std::forward<Args>(args)...);
                return result;
            }
            if (blocks.empty() || int(blocks.back().size()) == pool_block_size) {
                blocks.emplace_back();
                blocks.back().reserve(pool_block_size);
            }
            blocks.back().emplace_back(std::forward<Args>(args)...);
            return &blocks.back().back();
        }

        void recycle(Node* node) {
            free_nodes.push_back(node);
        }
    };

    inline static Pool pool;

    Node* root;
    Compare comp;

    static int subtree_size(const Node* t) {
        return t == nullptr ? 0 : t->size;
    }

    Node* new_node(T key) {
        return pool.emplace(std::move(key));
    }

    static void update(Node* t) {
        t->size = 1 + subtree_size(t->l) + subtree_size(t->r);
    }

    static Node* rotate_right(Node* t) {
        Node* s = t->l;
        t->l = s->r;
        s->r = t;
        update(t);
        update(s);
        return s;
    }

    static Node* rotate_left(Node* t) {
        Node* s = t->r;
        t->r = s->l;
        s->l = t;
        update(t);
        update(s);
        return s;
    }

    static Node* balance(Node* t) {
        if (t == nullptr) return nullptr;
        const int left_size = subtree_size(t->l);
        const int right_size = subtree_size(t->r);
        if (left_size + right_size > 1 && left_size > 3LL * right_size) {
            if (subtree_size(t->l->r) >= 2LL * subtree_size(t->l->l)) {
                t->l = rotate_left(t->l);
            }
            return rotate_right(t);
        }
        if (left_size + right_size > 1 && right_size > 3LL * left_size) {
            if (subtree_size(t->r->l) >= 2LL * subtree_size(t->r->r)) {
                t->r = rotate_right(t->r);
            }
            return rotate_left(t);
        }
        update(t);
        return t;
    }

    static Node* join_with_root(Node* l, Node* middle, Node* r) {
        const int left_size = subtree_size(l);
        const int right_size = subtree_size(r);
        if (left_size > 3LL * (right_size + 1)) {
            l->r = join_with_root(l->r, middle, r);
            return balance(l);
        }
        if (right_size > 3LL * (left_size + 1)) {
            r->l = join_with_root(l, middle, r->l);
            return balance(r);
        }
        middle->l = l;
        middle->r = r;
        return balance(middle);
    }

    static Node* detach_max(Node* t, Node*& maximum) {
        if (t->r == nullptr) {
            maximum = t;
            return t->l;
        }
        t->r = detach_max(t->r, maximum);
        return balance(t);
    }

    static Node* merge_nodes(Node* l, Node* r) {
        if (l == nullptr || r == nullptr) return l == nullptr ? r : l;
        Node* middle;
        l = detach_max(l, middle);
        return join_with_root(l, middle, r);
    }

    std::pair<Node*, Node*> split_nodes(Node* t, const T& key) {
        if (t == nullptr) return {nullptr, nullptr};
        Node* left = t->l;
        Node* right = t->r;
        t->l = nullptr;
        t->r = nullptr;
        if (comp(t->key, key)) {
            auto [l, r] = split_nodes(right, key);
            return {join_with_root(left, t, l), r};
        }
        auto [l, r] = split_nodes(left, key);
        return {l, join_with_root(r, t, right)};
    }

    Node* insert_impl(Node* t, T& key, bool& inserted) {
        if (t == nullptr) {
            inserted = true;
            return new_node(std::move(key));
        }
        if (comp(key, t->key)) {
            t->l = insert_impl(t->l, key, inserted);
        } else if (comp(t->key, key)) {
            t->r = insert_impl(t->r, key, inserted);
        } else {
            return t;
        }
        if (!inserted) return t;
        return balance(t);
    }

    Node* erase_impl(Node* t, const T& key, bool& erased) {
        if (t == nullptr) return nullptr;
        if (comp(key, t->key)) {
            t->l = erase_impl(t->l, key, erased);
        } else if (comp(t->key, key)) {
            t->r = erase_impl(t->r, key, erased);
        } else {
            erased = true;
            Node* l = t->l;
            Node* r = t->r;
            pool.recycle(t);
            return merge_nodes(l, r);
        }
        if (!erased) return t;
        return balance(t);
    }

    static const T* kth_impl(const Node* t, int k) {
        while (t != nullptr) {
            const int left_size = subtree_size(t->l);
            if (k < left_size) {
                t = t->l;
            } else if (k == left_size) {
                return &t->key;
            } else {
                k -= left_size + 1;
                t = t->r;
            }
        }
        return nullptr;
    }

    int order_of_key_impl(const Node* t, const T& key, bool upper) const {
        int result = 0;
        while (t != nullptr) {
            const bool take = upper ? !comp(key, t->key) : comp(t->key, key);
            if (take) {
                result += subtree_size(t->l) + 1;
                t = t->r;
            } else {
                t = t->l;
            }
        }
        return result;
    }

    const T* lower_bound_impl(const Node* t, const T& key, bool strict) const {
        const T* result = nullptr;
        while (t != nullptr) {
            const bool candidate = strict ? comp(key, t->key) : !comp(t->key, key);
            if (candidate) {
                result = &t->key;
                t = t->l;
            } else {
                t = t->r;
            }
        }
        return result;
    }

    const T* max_less_impl(const Node* t, const T& key, bool strict) const {
        const T* result = nullptr;
        while (t != nullptr) {
            const bool candidate = strict ? comp(t->key, key) : !comp(key, t->key);
            if (candidate) {
                result = &t->key;
                t = t->r;
            } else {
                t = t->l;
            }
        }
        return result;
    }

    bool contains_impl(const Node* t, const T& key) const {
        while (t != nullptr) {
            if (comp(key, t->key)) {
                t = t->l;
            } else if (comp(t->key, key)) {
                t = t->r;
            } else {
                return true;
            }
        }
        return false;
    }

    static void dump_impl(const Node* t, std::vector<T>& result) {
        if (t == nullptr) return;
        dump_impl(t->l, result);
        result.push_back(t->key);
        dump_impl(t->r, result);
    }

    static void recycle_impl(Node* t) {
        if (t == nullptr) return;
        recycle_impl(t->l);
        recycle_impl(t->r);
        pool.recycle(t);
    }

    Node* clone_impl(const Node* t) {
        if (t == nullptr) return nullptr;
        Node* result = new_node(t->key);
        result->l = clone_impl(t->l);
        result->r = clone_impl(t->r);
        update(result);
        return result;
    }

    OrderedSet(Node* node, Compare compare) : root(node), comp(std::move(compare)) {}

   public:
    explicit OrderedSet(Compare compare)
        : root(nullptr), comp(std::move(compare)) {}

    OrderedSet() : OrderedSet(Compare()) {}

    OrderedSet(std::initializer_list<T> init, Compare compare = Compare()) : OrderedSet(std::move(compare)) {
        for (const T& x : init) insert(x);
    }

    template <typename Iterator>
    OrderedSet(Iterator first, Iterator last, Compare compare = Compare()) : OrderedSet(std::move(compare)) {
        while (first != last) insert(*first++);
    }

    OrderedSet(const OrderedSet& other)
        : root(nullptr), comp(other.comp) {
        root = clone_impl(other.root);
    }

    OrderedSet(OrderedSet&& other) noexcept
        : root(std::exchange(other.root, nullptr)), comp(std::move(other.comp)) {}

    ~OrderedSet() {
        recycle_impl(root);
    }

    OrderedSet& operator=(OrderedSet other) {
        swap(other);
        return *this;
    }

    void swap(OrderedSet& other) noexcept {
        using std::swap;
        swap(root, other.root);
        swap(comp, other.comp);
    }

    int size() const { return subtree_size(root); }
    int unique_size() const { return size(); }
    bool empty() const { return root == nullptr; }

    void clear() {
        recycle_impl(root);
        root = nullptr;
    }

    bool insert(T key) {
        bool inserted = false;
        root = insert_impl(root, key, inserted);
        return inserted;
    }

    bool erase(const T& key) {
        bool erased = false;
        root = erase_impl(root, key, erased);
        return erased;
    }

    bool contains(const T& key) const { return contains_impl(root, key); }
    int count(const T& key) const { return contains(key) ? 1 : 0; }

    const T* find_by_order(int k) const {
        assert(0 <= k && k < size());
        return kth_impl(root, k);
    }

    T kth(int k) const { return *find_by_order(k); }
    int order_of_key(const T& key) const { return order_of_key_impl(root, key, false); }
    int count_less(const T& key) const { return order_of_key(key); }
    int count_less_equal(const T& key) const { return order_of_key_impl(root, key, true); }
    int count_greater(const T& key) const { return size() - count_less_equal(key); }
    int count_greater_equal(const T& key) const { return size() - count_less(key); }
    const T* lower_bound(const T& key) const { return lower_bound_impl(root, key, false); }
    const T* upper_bound(const T& key) const { return lower_bound_impl(root, key, true); }
    const T* min_ge(const T& key) const { return lower_bound(key); }
    const T* min_gt(const T& key) const { return upper_bound(key); }
    const T* max_le(const T& key) const { return max_less_impl(root, key, false); }
    const T* max_lt(const T& key) const { return max_less_impl(root, key, true); }
    const T* min() const { return empty() ? nullptr : kth_impl(root, 0); }
    const T* max() const { return empty() ? nullptr : kth_impl(root, size() - 1); }

    std::pair<OrderedSet, OrderedSet> split(const T& key) && {
        auto [l, r] = split_nodes(root, key);
        root = nullptr;
        return {OrderedSet(l, comp), OrderedSet(r, std::move(comp))};
    }

    OrderedSet merge(OrderedSet other) && {
        assert(empty() || other.empty() || comp(*max(), *other.min()));
        root = merge_nodes(root, other.root);
        other.root = nullptr;
        return std::move(*this);
    }

    std::vector<T> to_vector() const {
        std::vector<T> result;
        result.reserve(size());
        dump_impl(root, result);
        return result;
    }
};

}  // namespace ds
}  // namespace m1une


#line 12 "ds/bst/rollback_ordered_set.hpp"

namespace m1une {
namespace ds {

template <class T, class Compare = std::less<T>>
struct RollbackOrderedSet {
   private:
    enum class Kind { insert, erase, clear, merge };
    struct Entry {
        Kind kind;
        bool changed;
        std::optional<T> key;
        std::vector<T> keys;
    };

    OrderedSet<T, Compare> _data;
    std::vector<Entry> _history;
    std::vector<std::size_t> _checkpoints;

   public:
    explicit RollbackOrderedSet(Compare compare)
        : _data(std::move(compare)) {}
    RollbackOrderedSet() = default;

    RollbackOrderedSet(
        std::initializer_list<T> init,
        Compare compare = Compare()
    ) : _data(init, std::move(compare)) {}

    template <class Iterator>
    RollbackOrderedSet(
        Iterator first,
        Iterator last,
        Compare compare = Compare()
    ) : _data(first, last, std::move(compare)) {}

    int size() const { return _data.size(); }
    int unique_size() const { return _data.size(); }
    bool empty() const { return _data.empty(); }
    std::size_t node_count() const { return std::size_t(size()); }

    void clear() {
        if (_checkpoints.empty()) {
            _data.clear();
            return;
        }
        Entry entry{Kind::clear, !empty(), std::nullopt, {}};
        if (!empty()) entry.keys = _data.to_vector();
        _data.clear();
        _history.push_back(std::move(entry));
    }

    bool insert(T key) {
        if (_checkpoints.empty()) return _data.insert(std::move(key));
        bool changed = !_data.contains(key);
        Entry entry{Kind::insert, changed, std::nullopt, {}};
        if (changed) entry.key.emplace(key);
        _data.insert(std::move(key));
        _history.push_back(std::move(entry));
        return changed;
    }

    bool erase(const T& key) {
        if (_checkpoints.empty()) return _data.erase(key);
        bool changed = _data.contains(key);
        Entry entry{Kind::erase, changed, std::nullopt, {}};
        if (changed) entry.key.emplace(key);
        _data.erase(key);
        _history.push_back(std::move(entry));
        return changed;
    }

    void merge(const RollbackOrderedSet& other) {
        std::vector<T> keys = other.to_vector();
        for (const T& key : keys) {
            bool inserted = _data.insert(key);
            assert(inserted);
        }
        if (!_checkpoints.empty()) {
            _history.push_back(Entry{Kind::merge, !keys.empty(), std::nullopt, std::move(keys)});
        }
    }

    void merge(const OrderedSet<T, Compare>& other) {
        std::vector<T> keys = other.to_vector();
        for (const T& key : keys) {
            bool inserted = _data.insert(key);
            assert(inserted);
        }
        if (!_checkpoints.empty()) {
            _history.push_back(Entry{Kind::merge, !keys.empty(), std::nullopt, std::move(keys)});
        }
    }

    bool contains(const T& key) const { return _data.contains(key); }
    int count(const T& key) const { return _data.count(key); }
    const T* find_by_order(int order) const { return _data.find_by_order(order); }
    T kth(int order) const { return _data.kth(order); }
    int order_of_key(const T& key) const { return _data.order_of_key(key); }
    int count_less(const T& key) const { return _data.count_less(key); }
    int count_less_equal(const T& key) const { return _data.count_less_equal(key); }
    int count_greater(const T& key) const { return _data.count_greater(key); }
    int count_greater_equal(const T& key) const { return _data.count_greater_equal(key); }
    const T* lower_bound(const T& key) const { return _data.lower_bound(key); }
    const T* upper_bound(const T& key) const { return _data.upper_bound(key); }
    const T* min_ge(const T& key) const { return _data.min_ge(key); }
    const T* min_gt(const T& key) const { return _data.min_gt(key); }
    const T* max_le(const T& key) const { return _data.max_le(key); }
    const T* max_lt(const T& key) const { return _data.max_lt(key); }
    const T* min() const { return _data.min(); }
    const T* max() const { return _data.max(); }
    std::vector<T> to_vector() const { return _data.to_vector(); }

    int snapshot() { _checkpoints.push_back(_history.size()); return int(_checkpoints.size()); }
    int snapshot_count() const { return int(_checkpoints.size()); }

    void reserve_snapshots(int count) {
        assert(0 <= count);
        _checkpoints.reserve(count);
    }

   private:
    void restore_one() {
        Entry entry = std::move(_history.back());
        _history.pop_back();
        if (!entry.changed) return;
        if (entry.kind == Kind::insert) {
            _data.erase(*entry.key);
        } else if (entry.kind == Kind::erase) {
            _data.insert(std::move(*entry.key));
        } else if (entry.kind == Kind::clear) {
            for (T& key : entry.keys) _data.insert(std::move(key));
        } else {
            for (const T& key : entry.keys) _data.erase(key);
        }
    }

   public:

    void rollback(int state) {
        assert(1 <= state && state <= snapshot_count());
        while (_history.size() > _checkpoints[state - 1]) restore_one();
        _checkpoints.resize(state);
    }

    void clear_history() { _history.clear(); _checkpoints.clear(); }

    void release() {
        _data.clear();
        _history.clear();
        _checkpoints.clear();
    }
};

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