Rollback Ordered Set
(ds/bst/rollback_ordered_set.hpp)
- View this file on GitHub
- Last update: 2026-08-12 17:21:09+09:00
- Include:
#include "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