Skip to content
File

Blob: src/workerd/util/small-set.h

cpp398 lines
1// Copyright (c) 2017-2022 Cloudflare, Inc.
2// Licensed under the Apache 2.0 license found in the LICENSE file or at:
3// https://opensource.org/licenses/Apache-2.0
4 
5#pragma once
6 
7#include <kj/common.h>
8#include <kj/debug.h>
9#include <kj/one-of.h>
10#include <kj/vector.h>
11 
12#include <concepts>
13 
14namespace workerd {
15 
16// Concept for types that support reference counting via addRef().
17// This matches kj::Rc<T> which has an addRef() method returning kj::Rc<T>.
18template <typename T>
19concept RefCountedSmartPtr = requires(T& t) {
20 { t.addRef() } -> std::convertible_to<T>;
21};
22 
23// Concept for smart pointers to WeakRef-like types that have a tryGet() method.
24// This is used to constrain forEach() which needs to check if the ref is still valid.
25template <typename T>
26concept WeakRefSmartPtr = RefCountedSmartPtr<T> && requires(T& t) {
27 { t->tryGet() };
28};
29 
30// A set-like container optimized for the common case of storing 0-2 items
31// of reference-counted smart pointer types (like kj::Rc<T>).
32//
33// This uses a kj::OneOf to avoid heap allocations for small sets.
34//
35// Performance characteristics:
36// - 0-1 items: Zero heap allocations, O(1) operations
37// - 2 items: Zero heap allocations, O(1) operations
38// - 3+ items: Single heap allocation (kj::Vector), O(n) operations
39//
40// Typical usage patterns:
41// - 99% of instances have 1 item
42// - 0.9% of instances have 2 items
43// - 0.1% of instances have 3+ items
44//
45// This is NOT a drop-in replacement for std::set because:
46// - Items are not kept in sorted order
47// - No logarithmic lookup guarantees
48// - Optimized for small sizes only
49//
50// Iterator invalidation:
51// - Iterators are invalidated when items are removed or storage state changes
52// - If iterating over items that may be removed during iteration, use releaseSnapshot()
53// to get owned copies that remain valid even if the original is removed from the set.
54//
55// Template parameter T must be a reference-counted smart pointer type like kj::Rc<X>
56// that has an addRef() method returning the same type.
57template <RefCountedSmartPtr T>
58class SmallSet {
59 public:
60 SmallSet() = default;
61 KJ_DISALLOW_COPY(SmallSet);
62 SmallSet(SmallSet&&) = default;
63 SmallSet& operator=(SmallSet&&) = default;
64 
65 // Add an item to the set. The item is moved into the set.
66 // For move-only types, use containsIf() first to check for duplicates if needed.
67 void add(T item) {
68 KJ_SWITCH_ONEOF(storage) {
69 KJ_CASE_ONEOF(none, None) {
70 storage = Single(kj::mv(item));
71 return;
72 }
73 KJ_CASE_ONEOF(single, Single) {
74 storage = Double(kj::mv(single.item), kj::mv(item));
75 return;
76 }
77 KJ_CASE_ONEOF(dbl, Double) {
78 auto vec = kj::Vector<T>(4);
79 vec.add(kj::mv(dbl.first));
80 vec.add(kj::mv(dbl.second));
81 vec.add(kj::mv(item));
82 storage = kj::mv(vec);
83 return;
84 }
85 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
86 vec.add(kj::mv(item));
87 return;
88 }
89 }
90 KJ_UNREACHABLE;
91 }
92 
93 // Remove an item matching the predicate. Returns true if an item was removed.
94 // The predicate receives a const reference to each item.
95 template <typename Predicate>
96 bool removeIf(Predicate&& predicate) {
97 KJ_SWITCH_ONEOF(storage) {
98 KJ_CASE_ONEOF(none, None) {
99 return false;
100 }
101 KJ_CASE_ONEOF(single, Single) {
102 if (predicate(single.item)) {
103 storage = None();
104 return true;
105 }
106 return false;
107 }
108 KJ_CASE_ONEOF(dbl, Double) {
109 if (predicate(dbl.first)) {
110 storage = Single(kj::mv(dbl.second));
111 return true;
112 }
113 if (predicate(dbl.second)) {
114 storage = Single(kj::mv(dbl.first));
115 return true;
116 }
117 return false;
118 }
119 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
120 // Find and remove the first matching item
121 for (size_t i = 0; i < vec.size(); ++i) {
122 if (predicate(vec[i])) {
123 // Remove by overwriting with last element and truncating
124 if (i < vec.size() - 1) {
125 vec[i] = kj::mv(vec.back());
126 }
127 vec.removeLast();
128 
129 // Transition back to smaller state if appropriate
130 if (vec.size() == 2) {
131 storage = Double(kj::mv(vec[0]), kj::mv(vec[1]));
132 } else if (vec.size() == 1) {
133 storage = Single(kj::mv(vec[0]));
134 } else if (vec.size() == 0) {
135 storage = None();
136 }
137 // else: vec.size() >= 3, stay in Vector state
138 
139 return true;
140 }
141 }
142 return false;
143 }
144 }
145 KJ_UNREACHABLE;
146 }
147 
148 // Check if the set contains an item matching the predicate.
149 template <typename Predicate>
150 bool containsIf(Predicate&& predicate) const {
151 KJ_SWITCH_ONEOF(storage) {
152 KJ_CASE_ONEOF(none, None) {
153 return false;
154 }
155 KJ_CASE_ONEOF(single, Single) {
156 return predicate(single.item);
157 }
158 KJ_CASE_ONEOF(dbl, Double) {
159 return predicate(dbl.first) || predicate(dbl.second);
160 }
161 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
162 for (auto& existing: vec) {
163 if (predicate(existing)) return true;
164 }
165 return false;
166 }
167 }
168 KJ_UNREACHABLE;
169 }
170 
171 // Get the number of items in the set.
172 size_t size() const {
173 KJ_SWITCH_ONEOF(storage) {
174 KJ_CASE_ONEOF(none, None) {
175 return 0;
176 }
177 KJ_CASE_ONEOF(single, Single) {
178 return 1;
179 }
180 KJ_CASE_ONEOF(dbl, Double) {
181 return 2;
182 }
183 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
184 return vec.size();
185 }
186 }
187 KJ_UNREACHABLE;
188 }
189 
190 // Check if the set is empty.
191 bool empty() const {
192 return size() == 0;
193 }
194 
195 // Clear all items from the set.
196 void clear() {
197 storage = None();
198 }
199 
200 // Iterate over all valid (non-invalidated) WeakRef items, calling func for each.
201 // This is safe to use even if func modifies the set (e.g., removes items).
202 //
203 // Only available when T is a smart pointer to a WeakRef-like type with tryGet() method
204 // (e.g., kj::Rc<WeakRef<X>>). The callback receives a reference to the
205 // underlying type (X&).
206 //
207 // Example:
208 // SmallSet<kj::Rc<WeakRef<Consumer>>> consumers;
209 // consumers.forEach([&](Consumer& c) {
210 // c.close(js); // Safe even if this removes other consumers
211 // });
212 template <typename F>
213 void forEach(F&& func)
214 requires WeakRefSmartPtr<T>
215 {
216 KJ_SWITCH_ONEOF(storage) {
217 KJ_CASE_ONEOF(none, None) {
218 return;
219 }
220 KJ_CASE_ONEOF(single, Single) {
221 KJ_IF_SOME(ref, single.item->tryGet()) {
222 func(ref);
223 }
224 return;
225 }
226 KJ_CASE_ONEOF(dbl, Double) {
227 // The storage state may change during iteration if func modifies the set,
228 // so we take snapshots of the items first. Snapshotting just requires calling
229 // addRef to increment the ref counts.
230 kj::Array<T> refs = kj::arr(dbl.first.addRef(), dbl.second.addRef());
231 for (auto& item: refs) {
232 // We check tryGet on each item in case func invalidated some of them
233 // in prior iterations.
234 KJ_IF_SOME(ref, item->tryGet()) {
235 func(ref);
236 }
237 }
238 return;
239 }
240 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
241 // The storage state may change during iteration if func modifies the set,
242 // so we take snapshots of the items first. Snapshotting just requires calling
243 // addRef to increment the ref counts.
244 auto snapshot = KJ_MAP(item, vec) { return item.addRef(); };
245 for (auto& item: snapshot) {
246 // We check tryGet on each item in case func invalidated some of them
247 // in prior iterations.
248 KJ_IF_SOME(ref, item->tryGet()) {
249 func(ref);
250 }
251 }
252 return;
253 }
254 }
255 KJ_UNREACHABLE;
256 }
257 
258 private:
259 struct None {};
260 
261 struct Single {
262 T item;
263 explicit Single(T item): item(kj::mv(item)) {}
264 };
265 
266 struct Double {
267 T first;
268 T second;
269 Double(T first, T second): first(kj::mv(first)), second(kj::mv(second)) {}
270 };
271 
272 using Storage = kj::OneOf<None, Single, Double, kj::Vector<T>>;
273 Storage storage = None();
274 
275 public:
276 // Iterator support - returns const references to items
277 class ConstIterator {
278 public:
279 ConstIterator() = default;
280 
281 const T& operator*() const {
282 KJ_SWITCH_ONEOF(*storage) {
283 KJ_CASE_ONEOF(none, None) {
284 KJ_FAIL_REQUIRE("Dereferencing end iterator");
285 }
286 KJ_CASE_ONEOF(single, Single) {
287 KJ_REQUIRE(index == 0, "Invalid iterator");
288 return single.item;
289 }
290 KJ_CASE_ONEOF(dbl, Double) {
291 KJ_REQUIRE(index < 2, "Invalid iterator");
292 return index == 0 ? dbl.first : dbl.second;
293 }
294 KJ_CASE_ONEOF(vec, kj::Vector<T>) {
295 KJ_REQUIRE(index < vec.size(), "Invalid iterator");
296 return vec[index];
297 }
298 }
299 KJ_UNREACHABLE;
300 }
301 
302 ConstIterator& operator++() {
303 ++index;
304 return *this;
305 }
306 
307 ConstIterator operator++(int) {
308 ConstIterator tmp = *this;
309 ++index;
310 return tmp;
311 }
312 
313 bool operator==(const ConstIterator& other) const {
314 return storage == other.storage && index == other.index;
315 }
316 
317 bool operator!=(const ConstIterator& other) const {
318 return !(*this == other);
319 }
320 
321 private:
322 friend class SmallSet;
323 
324 ConstIterator(const Storage* storage, size_t index): storage(storage), index(index) {}
325 
326 const Storage* storage = nullptr;
327 size_t index = 0;
328 };
329 
330 // Mutable iterator - returns mutable references to items
331 class Iterator {
332 public:
333 Iterator() = default;
334 
335 T& operator*() const {
336 if (storage->template is<None>()) {
337 KJ_FAIL_REQUIRE("Dereferencing end iterator");
338 } else if (storage->template is<Single>()) {
339 KJ_REQUIRE(index == 0, "Invalid iterator");
340 return storage->template get<Single>().item;
341 } else if (storage->template is<Double>()) {
342 KJ_REQUIRE(index < 2, "Invalid iterator");
343 auto& dbl = storage->template get<Double>();
344 return index == 0 ? dbl.first : dbl.second;
345 } else {
346 auto& vec = storage->template get<kj::Vector<T>>();
347 KJ_REQUIRE(index < vec.size(), "Invalid iterator");
348 return vec[index];
349 }
350 }
351 
352 Iterator& operator++() {
353 ++index;
354 return *this;
355 }
356 
357 Iterator operator++(int) {
358 Iterator tmp = *this;
359 ++index;
360 return tmp;
361 }
362 
363 bool operator==(const Iterator& other) const {
364 return storage == other.storage && index == other.index;
365 }
366 
367 bool operator!=(const Iterator& other) const {
368 return !(*this == other);
369 }
370 
371 private:
372 friend class SmallSet;
373 
374 Iterator(Storage* storage, size_t index): storage(storage), index(index) {}
375 
376 Storage* storage = nullptr;
377 size_t index = 0;
378 };
379 
380 Iterator begin() {
381 return Iterator(&storage, 0);
382 }
383 
384 Iterator end() {
385 return Iterator(&storage, size());
386 }
387 
388 ConstIterator begin() const {
389 return ConstIterator(&storage, 0);
390 }
391 
392 ConstIterator end() const {
393 return ConstIterator(&storage, size());
394 }
395};
396 
397} // namespace workerd