Rollback Dynamic Segment Tree
(ds/segtree/rollback_dynamic_segtree.hpp)
- View this file on GitHub
- Last update: 2026-08-12 17:21:09+09:00
- Include:
#include "ds/segtree/rollback_dynamic_segtree.hpp"
Overview
RollbackDynamicSegtree<Monoid, Index> is a sparse segment tree over an
integral half-open domain with point assignment and rollback. Unvisited ranges
retain the configured initial value.
Methods
Constructors and read-only methods follow DynamicSegtree<Monoid, Index>.
| Method | Description | Complexity |
|---|---|---|
void set(Index pos, T value), void set_inplace(Index pos, T value)
|
Assigns one point. | $O(\log U)$ |
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) |
Restores a current-path snapshot. | $O(F)$ total |
void clear_history(), void release()
|
Releases saved states, or all materialized nodes. | $O(F)$ |
$U$ is the domain width and $F = O(\log U)$ per assignment.
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.
Within one snapshot interval, a materialized node is saved only before its first mutation; newly allocated nodes are truncated directly by rollback.
Example
#include "ds/segtree/rollback_dynamic_segtree.hpp"
#include "monoid/add.hpp"
using Add = m1une::monoid::Add<long long>;
m1une::ds::RollbackDynamicSegtree<Add> seg(-100, 100);
int state = seg.snapshot();
seg.set(-4, 7);
seg.rollback(state);
assert(seg.all_prod() == 0);
Depends on
ds/detail/rollback_journal.hpp
ds/segtree/dynamic_segtree_common.hpp
Monoid Concept
(monoid/concept.hpp)
Verified with
Code
#ifndef M1UNE_DS_SEGTREE_ROLLBACK_DYNAMIC_SEGTREE_HPP
#define M1UNE_DS_SEGTREE_ROLLBACK_DYNAMIC_SEGTREE_HPP 1
#include <array>
#include <cassert>
#include <concepts>
#include <limits>
#include <numeric>
#include <type_traits>
#include <utility>
#include "../../monoid/concept.hpp"
#include "../detail/rollback_journal.hpp"
#include "dynamic_segtree_common.hpp"
namespace m1une {
namespace ds {
template <m1une::monoid::IsMonoid Monoid, std::integral Index = long long>
requires(!std::same_as<std::remove_cv_t<Index>, bool>)
struct RollbackDynamicSegtree {
using T = typename Monoid::value_type;
using index_type = Index;
using size_type = detail::dynamic_size_type<Index>;
private:
struct Node {
T value = Monoid::id();
int left = 0;
int right = 0;
};
static constexpr int path_capacity = std::numeric_limits<size_type>::digits + 1;
detail::UniformMonoidDomain<Monoid, Index> _domain;
detail::RollbackJournal<Node> _journal;
int root() const { return _journal[0].left; }
int new_node() { return _journal.emplace(); }
const T& value(int node, Index left, Index right, int depth) const {
if (node) return _journal[node].value;
return _domain.default_product(depth, left, right);
}
void update(int node, Index left, Index right, int depth) {
Index middle = std::midpoint(left, right);
_journal.touch(node);
_journal[node].value = Monoid::op(
value(_journal[node].left, left, middle, depth + 1),
value(_journal[node].right, middle, right, depth + 1)
);
}
T prod_node(int node, Index left, Index right, int depth, Index query_left, Index query_right) const {
if (query_right <= left || right <= query_left) return Monoid::id();
if (query_left <= left && right <= query_right) return value(node, left, right, depth);
Index middle = std::midpoint(left, right);
return Monoid::op(
prod_node(node ? _journal[node].left : 0, left, middle, depth + 1, query_left, query_right),
prod_node(node ? _journal[node].right : 0, middle, right, depth + 1, query_left, query_right)
);
}
template <class Predicate>
Index max_right_node(int node, Index left, Index right, int depth, Index query_left, T& product,
Predicate& predicate) const {
if (right <= query_left) return right;
if (query_left <= left) {
T next = Monoid::op(product, value(node, left, right, depth));
if (predicate(next)) {
product = std::move(next);
return right;
}
Index middle = std::midpoint(left, right);
if (middle == left) return left;
}
Index middle = std::midpoint(left, right);
Index result = max_right_node(node ? _journal[node].left : 0, left, middle, depth + 1,
query_left, product, predicate);
if (result < middle) return result;
return max_right_node(node ? _journal[node].right : 0, middle, right, depth + 1,
query_left, product, predicate);
}
template <class Predicate>
Index min_left_node(int node, Index left, Index right, int depth, Index query_right, T& product,
Predicate& predicate) const {
if (query_right <= left) return left;
if (right <= query_right) {
T next = Monoid::op(value(node, left, right, depth), product);
if (predicate(next)) {
product = std::move(next);
return left;
}
Index middle = std::midpoint(left, right);
if (middle == left) return right;
}
Index middle = std::midpoint(left, right);
Index result = min_left_node(node ? _journal[node].right : 0, middle, right, depth + 1,
query_right, product, predicate);
if (middle < result) return result;
return min_left_node(node ? _journal[node].left : 0, left, middle, depth + 1,
query_right, product, predicate);
}
public:
RollbackDynamicSegtree() : RollbackDynamicSegtree(Index(0), Index(0)) {}
explicit RollbackDynamicSegtree(Index n) : RollbackDynamicSegtree(Index(0), n) {
if constexpr (std::signed_integral<Index>) assert(Index(0) <= n);
}
RollbackDynamicSegtree(Index left, Index right)
: RollbackDynamicSegtree(left, right, Monoid::id()) {}
RollbackDynamicSegtree(Index left, Index right, T initial_value)
: _domain(left, right, std::move(initial_value)) {
_journal.emplace();
}
size_type size() const { return _domain.size(); }
bool empty() const { return _domain.empty(); }
Index left_bound() const { return _domain.left_bound(); }
Index right_bound() const { return _domain.right_bound(); }
const T& initial_value() const { return _domain.initial_value(); }
void reserve(std::size_t node_capacity) {
_journal.nodes.reserve(node_capacity + 1);
_journal.saved_epoch.reserve(node_capacity + 1);
}
std::size_t node_count() const { return _journal.nodes.size() - 1; }
void set(Index pos, T x) {
assert(left_bound() <= pos && pos < right_bound());
if (!root()) {
int node = new_node();
_journal.touch(0);
_journal[0].left = node;
}
std::array<int, path_capacity> path;
std::array<Index, path_capacity> path_left;
std::array<Index, path_capacity> path_right;
int depth = 0;
int node = root();
Index left = left_bound();
Index right = right_bound();
while (true) {
path[depth] = node;
path_left[depth] = left;
path_right[depth] = right;
++depth;
Index middle = std::midpoint(left, right);
if (middle == left) break;
if (pos < middle) {
if (!_journal[node].left) {
int child = new_node();
_journal.touch(node);
_journal[node].left = child;
}
node = _journal[node].left;
right = middle;
} else {
if (!_journal[node].right) {
int child = new_node();
_journal.touch(node);
_journal[node].right = child;
}
node = _journal[node].right;
left = middle;
}
}
_journal.touch(node);
_journal[node].value = std::move(x);
for (int index = depth - 2; index >= 0; --index) {
update(path[index], path_left[index], path_right[index], index);
}
}
void set_inplace(Index pos, T x) { set(pos, std::move(x)); }
T get(Index pos) const {
assert(left_bound() <= pos && pos < right_bound());
int node = root();
Index left = left_bound();
Index right = right_bound();
int depth = 0;
while (node) {
Index middle = std::midpoint(left, right);
if (middle == left) return value(node, left, right, depth);
if (pos < middle) {
node = _journal[node].left;
right = middle;
} else {
node = _journal[node].right;
left = middle;
}
++depth;
}
return initial_value();
}
T operator[](Index pos) const { return get(pos); }
T prod(Index left, Index right) const {
assert(left_bound() <= left && left <= right && right <= right_bound());
if (left == right) return Monoid::id();
return prod_node(root(), left_bound(), right_bound(), 0, left, right);
}
T all_prod() const { return value(root(), left_bound(), right_bound(), 0); }
template <class Predicate>
Index max_right(Index left, Predicate predicate) const {
assert(left_bound() <= left && left <= right_bound());
assert(predicate(Monoid::id()));
if (left == right_bound()) return right_bound();
T product = Monoid::id();
return max_right_node(root(), left_bound(), right_bound(), 0, left, product, predicate);
}
template <class Predicate>
Index min_left(Index right, Predicate predicate) const {
assert(left_bound() <= right && right <= right_bound());
assert(predicate(Monoid::id()));
if (right == left_bound()) return left_bound();
T product = Monoid::id();
return min_left_node(root(), left_bound(), right_bound(), 0, right, product, predicate);
}
int snapshot() { return _journal.snapshot(); }
int snapshot_count() const { return _journal.snapshot_count(); }
void reserve_snapshots(int count) { _journal.reserve_snapshots(count); }
void rollback(int state) { _journal.rollback(state); }
void clear_history() { _journal.clear_history(); }
void release() { _journal.clear(); _journal.emplace(); }
};
} // namespace ds
} // namespace m1une
#endif // M1UNE_DS_SEGTREE_ROLLBACK_DYNAMIC_SEGTREE_HPP#line 1 "ds/segtree/rollback_dynamic_segtree.hpp"
#include <array>
#include <cassert>
#include <concepts>
#include <limits>
#include <numeric>
#include <type_traits>
#include <utility>
#line 1 "monoid/concept.hpp"
#line 5 "monoid/concept.hpp"
namespace m1une {
namespace monoid {
// Concept to check if a type satisfies the requirements of a Monoid.
// A Monoid must have a `value_type`, an identity element `id()`, and an associative binary operation `op()`.
template <typename M>
concept IsMonoid = requires(typename M::value_type a, typename M::value_type b) {
// 1. Must define `value_type`
typename M::value_type;
// 2. Must have a static method `id()` returning `value_type`
{ M::id() } -> std::same_as<typename M::value_type>;
// 3. Must have a static method `op(a, b)` returning `value_type`
{ M::op(a, b) } -> std::same_as<typename M::value_type>;
};
// Concept for groups. A type satisfying this concept must also obey the group
// laws; concepts can check the interface but not the algebraic properties.
template <typename M>
concept IsGroup = IsMonoid<M> && requires(typename M::value_type a) {
{ M::inv(a) } -> std::same_as<typename M::value_type>;
};
// Concept for commutative groups. Commutativity is a semantic requirement and
// cannot be checked by a C++ concept.
template <typename M>
concept IsCommutativeGroup = IsGroup<M>;
} // namespace monoid
} // namespace m1une
#line 1 "ds/detail/rollback_journal.hpp"
#include <algorithm>
#line 6 "ds/detail/rollback_journal.hpp"
#include <cstddef>
#include <cstdint>
#line 10 "ds/detail/rollback_journal.hpp"
#include <vector>
namespace m1une {
namespace ds {
namespace detail {
template <class Node>
struct RollbackJournal {
struct Change {
int index;
Node value;
};
struct Checkpoint {
std::size_t change_size;
std::size_t node_size;
std::uint64_t epoch;
};
std::vector<Node> nodes;
std::vector<Change> changes;
std::vector<Checkpoint> checkpoints;
std::vector<std::uint64_t> saved_epoch;
std::uint64_t next_epoch = 1;
std::uint64_t new_epoch() {
if (next_epoch == 0) {
std::fill(saved_epoch.begin(), saved_epoch.end(), 0);
next_epoch = 1;
}
return next_epoch++;
}
int size() const { return int(nodes.size()); }
Node& operator[](int index) { return nodes[index]; }
const Node& operator[](int index) const { return nodes[index]; }
template <class... Args>
int emplace(Args&&... args) {
assert(nodes.size() < std::size_t(std::numeric_limits<int>::max()));
int index = int(nodes.size());
nodes.emplace_back(std::forward<Args>(args)...);
saved_epoch.push_back(0);
return index;
}
int snapshot() {
assert(checkpoints.size() < std::size_t(std::numeric_limits<int>::max()));
checkpoints.push_back(Checkpoint{changes.size(), nodes.size(), new_epoch()});
return int(checkpoints.size());
}
void touch(int index) {
assert(0 <= index && index < size());
if (checkpoints.empty()) return;
const Checkpoint& checkpoint = checkpoints.back();
if (std::size_t(index) >= checkpoint.node_size) return;
if (saved_epoch[index] == checkpoint.epoch) return;
saved_epoch[index] = checkpoint.epoch;
changes.push_back(Change{index, nodes[index]});
}
int snapshot_count() const { return int(checkpoints.size()); }
void reserve_snapshots(int count) {
assert(0 <= count);
checkpoints.reserve(count);
}
void reserve_changes(std::size_t count) { changes.reserve(count); }
void rollback(int state) {
assert(1 <= state && state <= snapshot_count());
Checkpoint checkpoint = checkpoints[state - 1];
while (changes.size() > checkpoint.change_size) {
Change change = std::move(changes.back());
changes.pop_back();
nodes[change.index] = std::move(change.value);
}
nodes.erase(nodes.begin() + checkpoint.node_size, nodes.end());
saved_epoch.resize(checkpoint.node_size);
checkpoints.resize(state);
checkpoints.back().change_size = changes.size();
checkpoints.back().node_size = nodes.size();
checkpoints.back().epoch = new_epoch();
}
void clear_history() {
changes.clear();
checkpoints.clear();
std::fill(saved_epoch.begin(), saved_epoch.end(), 0);
}
void clear() {
nodes.clear();
changes.clear();
checkpoints.clear();
saved_epoch.clear();
next_epoch = 1;
}
};
} // namespace detail
} // namespace ds
} // namespace m1une
#line 1 "ds/segtree/dynamic_segtree_common.hpp"
#line 11 "ds/segtree/dynamic_segtree_common.hpp"
namespace m1une {
namespace ds {
namespace detail {
template <std::integral Index>
using dynamic_size_type = std::make_unsigned_t<Index>;
template <std::integral Index>
constexpr dynamic_size_type<Index> dynamic_distance(Index left, Index right) {
return static_cast<dynamic_size_type<Index>>(right) - static_cast<dynamic_size_type<Index>>(left);
}
template <class Monoid, class Size>
typename Monoid::value_type monoid_repeat(typename Monoid::value_type value, Size count) {
typename Monoid::value_type result = Monoid::id();
while (count != 0) {
if (count & 1) result = Monoid::op(result, value);
count >>= 1;
if (count != 0) value = Monoid::op(value, value);
}
return result;
}
template <class ActedMonoid>
typename ActedMonoid::value_type dynamic_mapping(
const typename ActedMonoid::operator_type& f,
const typename ActedMonoid::value_type& value
) {
using F = typename ActedMonoid::operator_type;
using T = typename ActedMonoid::value_type;
if constexpr (requires(F g, T x, long long ord) { ActedMonoid::mapping(g, x, ord); }) {
return ActedMonoid::mapping(f, value, 0);
} else {
return ActedMonoid::mapping(f, value);
}
}
template <class ActedMonoid, class Size>
typename ActedMonoid::operator_type dynamic_shift(
const typename ActedMonoid::operator_type& f,
Size offset
) {
using F = typename ActedMonoid::operator_type;
if constexpr (requires(F g, long long ord) { ActedMonoid::op_shift(g, ord); }) {
assert(offset <= static_cast<Size>(std::numeric_limits<long long>::max()));
return ActedMonoid::op_shift(f, static_cast<long long>(offset));
} else {
return f;
}
}
template <class Monoid, std::integral Index>
class UniformMonoidDomain {
public:
using T = typename Monoid::value_type;
using size_type = dynamic_size_type<Index>;
private:
struct Level {
size_type small_length;
T small_value;
T large_value;
};
Index _left;
Index _right;
T _initial_value;
std::vector<Level> _levels;
public:
UniformMonoidDomain(Index left, Index right, T initial_value)
: _left(left), _right(right), _initial_value(std::move(initial_value)) {
assert(left <= right);
size_type n = size();
constexpr int digits = std::numeric_limits<size_type>::digits;
_levels.reserve(digits + 1);
for (int depth = 0; depth <= digits; depth++) {
size_type small = depth == digits ? 0 : n >> depth;
size_type large = small;
if (depth != 0) {
bool has_remainder;
if (depth == digits) {
has_remainder = n != 0;
} else {
size_type mask = (size_type(1) << depth) - 1;
has_remainder = (n & mask) != 0;
}
if (has_remainder) large++;
}
_levels.push_back(Level{
small,
monoid_repeat<Monoid>(_initial_value, small),
monoid_repeat<Monoid>(_initial_value, large),
});
}
}
Index left_bound() const {
return _left;
}
Index right_bound() const {
return _right;
}
size_type size() const {
return dynamic_distance(_left, _right);
}
bool empty() const {
return _left == _right;
}
const T& initial_value() const {
return _initial_value;
}
const T& default_product(int depth, Index left, Index right) const {
assert(0 <= depth && depth < int(_levels.size()));
const Level& level = _levels[depth];
size_type length = dynamic_distance(left, right);
if (length == level.small_length) return level.small_value;
assert(length == level.small_length + 1);
return level.large_value;
}
};
} // namespace detail
} // namespace ds
} // namespace m1une
#line 15 "ds/segtree/rollback_dynamic_segtree.hpp"
namespace m1une {
namespace ds {
template <m1une::monoid::IsMonoid Monoid, std::integral Index = long long>
requires(!std::same_as<std::remove_cv_t<Index>, bool>)
struct RollbackDynamicSegtree {
using T = typename Monoid::value_type;
using index_type = Index;
using size_type = detail::dynamic_size_type<Index>;
private:
struct Node {
T value = Monoid::id();
int left = 0;
int right = 0;
};
static constexpr int path_capacity = std::numeric_limits<size_type>::digits + 1;
detail::UniformMonoidDomain<Monoid, Index> _domain;
detail::RollbackJournal<Node> _journal;
int root() const { return _journal[0].left; }
int new_node() { return _journal.emplace(); }
const T& value(int node, Index left, Index right, int depth) const {
if (node) return _journal[node].value;
return _domain.default_product(depth, left, right);
}
void update(int node, Index left, Index right, int depth) {
Index middle = std::midpoint(left, right);
_journal.touch(node);
_journal[node].value = Monoid::op(
value(_journal[node].left, left, middle, depth + 1),
value(_journal[node].right, middle, right, depth + 1)
);
}
T prod_node(int node, Index left, Index right, int depth, Index query_left, Index query_right) const {
if (query_right <= left || right <= query_left) return Monoid::id();
if (query_left <= left && right <= query_right) return value(node, left, right, depth);
Index middle = std::midpoint(left, right);
return Monoid::op(
prod_node(node ? _journal[node].left : 0, left, middle, depth + 1, query_left, query_right),
prod_node(node ? _journal[node].right : 0, middle, right, depth + 1, query_left, query_right)
);
}
template <class Predicate>
Index max_right_node(int node, Index left, Index right, int depth, Index query_left, T& product,
Predicate& predicate) const {
if (right <= query_left) return right;
if (query_left <= left) {
T next = Monoid::op(product, value(node, left, right, depth));
if (predicate(next)) {
product = std::move(next);
return right;
}
Index middle = std::midpoint(left, right);
if (middle == left) return left;
}
Index middle = std::midpoint(left, right);
Index result = max_right_node(node ? _journal[node].left : 0, left, middle, depth + 1,
query_left, product, predicate);
if (result < middle) return result;
return max_right_node(node ? _journal[node].right : 0, middle, right, depth + 1,
query_left, product, predicate);
}
template <class Predicate>
Index min_left_node(int node, Index left, Index right, int depth, Index query_right, T& product,
Predicate& predicate) const {
if (query_right <= left) return left;
if (right <= query_right) {
T next = Monoid::op(value(node, left, right, depth), product);
if (predicate(next)) {
product = std::move(next);
return left;
}
Index middle = std::midpoint(left, right);
if (middle == left) return right;
}
Index middle = std::midpoint(left, right);
Index result = min_left_node(node ? _journal[node].right : 0, middle, right, depth + 1,
query_right, product, predicate);
if (middle < result) return result;
return min_left_node(node ? _journal[node].left : 0, left, middle, depth + 1,
query_right, product, predicate);
}
public:
RollbackDynamicSegtree() : RollbackDynamicSegtree(Index(0), Index(0)) {}
explicit RollbackDynamicSegtree(Index n) : RollbackDynamicSegtree(Index(0), n) {
if constexpr (std::signed_integral<Index>) assert(Index(0) <= n);
}
RollbackDynamicSegtree(Index left, Index right)
: RollbackDynamicSegtree(left, right, Monoid::id()) {}
RollbackDynamicSegtree(Index left, Index right, T initial_value)
: _domain(left, right, std::move(initial_value)) {
_journal.emplace();
}
size_type size() const { return _domain.size(); }
bool empty() const { return _domain.empty(); }
Index left_bound() const { return _domain.left_bound(); }
Index right_bound() const { return _domain.right_bound(); }
const T& initial_value() const { return _domain.initial_value(); }
void reserve(std::size_t node_capacity) {
_journal.nodes.reserve(node_capacity + 1);
_journal.saved_epoch.reserve(node_capacity + 1);
}
std::size_t node_count() const { return _journal.nodes.size() - 1; }
void set(Index pos, T x) {
assert(left_bound() <= pos && pos < right_bound());
if (!root()) {
int node = new_node();
_journal.touch(0);
_journal[0].left = node;
}
std::array<int, path_capacity> path;
std::array<Index, path_capacity> path_left;
std::array<Index, path_capacity> path_right;
int depth = 0;
int node = root();
Index left = left_bound();
Index right = right_bound();
while (true) {
path[depth] = node;
path_left[depth] = left;
path_right[depth] = right;
++depth;
Index middle = std::midpoint(left, right);
if (middle == left) break;
if (pos < middle) {
if (!_journal[node].left) {
int child = new_node();
_journal.touch(node);
_journal[node].left = child;
}
node = _journal[node].left;
right = middle;
} else {
if (!_journal[node].right) {
int child = new_node();
_journal.touch(node);
_journal[node].right = child;
}
node = _journal[node].right;
left = middle;
}
}
_journal.touch(node);
_journal[node].value = std::move(x);
for (int index = depth - 2; index >= 0; --index) {
update(path[index], path_left[index], path_right[index], index);
}
}
void set_inplace(Index pos, T x) { set(pos, std::move(x)); }
T get(Index pos) const {
assert(left_bound() <= pos && pos < right_bound());
int node = root();
Index left = left_bound();
Index right = right_bound();
int depth = 0;
while (node) {
Index middle = std::midpoint(left, right);
if (middle == left) return value(node, left, right, depth);
if (pos < middle) {
node = _journal[node].left;
right = middle;
} else {
node = _journal[node].right;
left = middle;
}
++depth;
}
return initial_value();
}
T operator[](Index pos) const { return get(pos); }
T prod(Index left, Index right) const {
assert(left_bound() <= left && left <= right && right <= right_bound());
if (left == right) return Monoid::id();
return prod_node(root(), left_bound(), right_bound(), 0, left, right);
}
T all_prod() const { return value(root(), left_bound(), right_bound(), 0); }
template <class Predicate>
Index max_right(Index left, Predicate predicate) const {
assert(left_bound() <= left && left <= right_bound());
assert(predicate(Monoid::id()));
if (left == right_bound()) return right_bound();
T product = Monoid::id();
return max_right_node(root(), left_bound(), right_bound(), 0, left, product, predicate);
}
template <class Predicate>
Index min_left(Index right, Predicate predicate) const {
assert(left_bound() <= right && right <= right_bound());
assert(predicate(Monoid::id()));
if (right == left_bound()) return left_bound();
T product = Monoid::id();
return min_left_node(root(), left_bound(), right_bound(), 0, right, product, predicate);
}
int snapshot() { return _journal.snapshot(); }
int snapshot_count() const { return _journal.snapshot_count(); }
void reserve_snapshots(int count) { _journal.reserve_snapshots(count); }
void rollback(int state) { _journal.rollback(state); }
void clear_history() { _journal.clear_history(); }
void release() { _journal.clear(); _journal.emplace(); }
};
} // namespace ds
} // namespace m1une