commonware_utils/
priority_set.rs1use std::{
2 cmp::Ordering,
3 collections::{BTreeSet, HashMap, HashSet},
4 hash::Hash,
5};
6
7#[derive(Eq, PartialEq)]
9struct Entry<I: Ord + Hash + Clone, P: Ord + Copy> {
10 item: I,
11 priority: P,
12}
13
14impl<I: Ord + Hash + Clone, P: Ord + Copy> Ord for Entry<I, P> {
15 fn cmp(&self, other: &Self) -> Ordering {
16 match self.priority.cmp(&other.priority) {
17 Ordering::Equal => self.item.cmp(&other.item),
18 other => other,
19 }
20 }
21}
22
23impl<I: Ord + Hash + Clone, V: Ord + Copy> PartialOrd for Entry<I, V> {
24 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
25 Some(self.cmp(other))
26 }
27}
28
29pub struct PrioritySet<I: Ord + Hash + Clone, P: Ord + Copy> {
32 entries: BTreeSet<Entry<I, P>>,
33 keys: HashMap<I, P>,
34}
35
36impl<I: Ord + Hash + Clone, P: Ord + Copy> PrioritySet<I, P> {
37 #[allow(clippy::new_without_default)]
41 pub fn new() -> Self {
42 Self {
43 entries: BTreeSet::new(),
44 keys: HashMap::new(),
45 }
46 }
47
48 pub fn put(&mut self, item: I, priority: P) {
50 let entry = if let Some(old_priority) = self.keys.remove(&item) {
52 let mut old_entry = Entry {
54 item: item.clone(),
55 priority: old_priority,
56 };
57 self.entries.remove(&old_entry);
58
59 old_entry.priority = priority;
61 old_entry
62 } else {
63 Entry { item, priority }
64 };
65
66 self.keys.insert(entry.item.clone(), entry.priority);
68 self.entries.insert(entry);
69 }
70
71 pub fn get(&self, item: &I) -> Option<P> {
73 self.keys.get(item).cloned()
74 }
75
76 pub fn remove(&mut self, item: &I) -> bool {
80 let Some(entry) = self.keys.remove(item).map(|priority| Entry {
81 item: item.clone(),
82 priority,
83 }) else {
84 return false;
85 };
86 assert!(self.entries.remove(&entry));
87 true
88 }
89
90 pub fn reconcile(&mut self, keep: &[I], default: P) {
93 let mut retained: HashSet<_> = keep.iter().collect();
95 let to_remove = self
96 .keys
97 .keys()
98 .filter(|item| !retained.remove(*item))
99 .cloned()
100 .collect::<Vec<_>>();
101 for item in to_remove {
102 let priority = self.keys.remove(&item).unwrap();
103 let entry = Entry { item, priority };
104 self.entries.remove(&entry);
105 }
106
107 for item in retained {
109 self.put(item.clone(), default);
110 }
111 }
112
113 pub fn retain(&mut self, predicate: impl Fn(&I) -> bool) {
115 self.entries.retain(|entry| predicate(&entry.item));
116 self.keys.retain(|key, _| predicate(key));
117 }
118
119 pub fn contains(&self, item: &I) -> bool {
121 self.keys.contains_key(item)
122 }
123
124 pub fn peek(&self) -> Option<(&I, &P)> {
126 self.entries
127 .iter()
128 .next()
129 .map(|entry| (&entry.item, &entry.priority))
130 }
131
132 pub fn pop(&mut self) -> Option<(I, P)> {
134 self.entries.pop_first().map(|entry| {
135 self.keys.remove(&entry.item);
136 (entry.item, entry.priority)
137 })
138 }
139
140 pub fn clear(&mut self) {
142 self.entries.clear();
143 self.keys.clear();
144 }
145
146 pub fn iter(&self) -> impl Iterator<Item = (&I, &P)> {
148 self.entries
149 .iter()
150 .map(|entry| (&entry.item, &entry.priority))
151 }
152
153 pub fn len(&self) -> usize {
155 self.entries.len()
156 }
157
158 pub fn is_empty(&self) -> bool {
160 self.entries.is_empty()
161 }
162}
163
164#[cfg(test)]
165mod tests {
166 use super::*;
167 use std::time::Duration;
168
169 #[test]
170 fn test_put_remove_and_iter() {
171 let mut pq = PrioritySet::new();
173
174 let key1 = "key1";
176 let key2 = "key2";
177 pq.put(key1, Duration::from_secs(10));
178 pq.put(key2, Duration::from_secs(5));
179
180 let entries: Vec<_> = pq.iter().collect();
182 assert_eq!(entries.len(), 2);
183 assert_eq!(*entries[0].0, key2);
184 assert_eq!(*entries[1].0, key1);
185
186 pq.remove(&key1);
188
189 let entries: Vec<_> = pq.iter().collect();
191 assert_eq!(entries.len(), 1);
192 assert_eq!(*entries[0].0, key2);
193
194 pq.remove(&key1);
196
197 let entries: Vec<_> = pq.iter().collect();
199 assert_eq!(entries.len(), 1);
200 assert_eq!(*entries[0].0, key2);
201 }
202
203 #[test]
204 fn test_update() {
205 let mut pq = PrioritySet::new();
207
208 let key = "key";
210 pq.put(key, Duration::from_secs(10));
211 assert_eq!(pq.get(&key).unwrap(), Duration::from_secs(10));
212
213 pq.put(key, Duration::from_secs(5));
215 assert_eq!(pq.get(&key).unwrap(), Duration::from_secs(5));
216
217 let entries: Vec<_> = pq.iter().collect();
219 assert_eq!(entries.len(), 1);
220 assert_eq!(*entries[0].1, Duration::from_secs(5));
221 }
222
223 #[test]
224 fn test_reconcile() {
225 let mut pq = PrioritySet::new();
227
228 let key1 = "key1";
230 let key2 = "key2";
231 pq.put(key1, Duration::from_secs(10));
232 pq.put(key2, Duration::from_secs(5));
233
234 let key3 = "key3";
236 pq.reconcile(&[key1, key3], Duration::from_secs(2));
237
238 let entries: Vec<_> = pq.iter().collect();
240 assert_eq!(entries.len(), 2);
241 assert!(
242 entries
243 .iter()
244 .any(|e| *e.0 == key1 && *e.1 == Duration::from_secs(10))
245 );
246 assert!(
247 entries
248 .iter()
249 .any(|e| *e.0 == key3 && *e.1 == Duration::from_secs(2))
250 );
251 }
252
253 #[test]
254 fn test_retain() {
255 let mut pq = PrioritySet::new();
257
258 pq.put("key1", Duration::from_secs(10));
260 pq.put("key2", Duration::from_secs(5));
261 pq.put("item3", Duration::from_secs(15));
262
263 pq.retain(|key| key.starts_with("key"));
265
266 assert_eq!(pq.len(), 2);
268 assert!(pq.contains(&"key1"));
269 assert!(pq.contains(&"key2"));
270 assert!(!pq.contains(&"item3"));
271
272 let entries: Vec<_> = pq.iter().collect();
274 assert_eq!(entries.len(), 2);
275 assert_eq!(*entries[0].0, "key2");
276 assert_eq!(*entries[1].0, "key1");
277 }
278
279 #[test]
280 fn test_clear() {
281 let mut pq = PrioritySet::new();
283
284 pq.put("key1", Duration::from_secs(10));
286 pq.put("key2", Duration::from_secs(5));
287
288 pq.clear();
290
291 assert_eq!(pq.len(), 0);
293 assert!(pq.is_empty());
294 assert!(pq.iter().next().is_none());
295 assert!(pq.peek().is_none());
296 }
297}