hermes_core/segment/
tracker.rs1use std::collections::{HashMap, HashSet};
11use std::sync::Arc;
12
13use parking_lot::Mutex;
14
15use crate::segment::SegmentId;
16use crate::segment::TrainedVectorStructures;
17
18struct TrackerInner {
20 ref_counts: HashMap<String, usize>,
21 pending_deletions: HashMap<String, SegmentId>,
22 scheduled_deletions: HashSet<String>,
24}
25
26pub struct SegmentTracker {
28 inner: Mutex<TrackerInner>,
29}
30
31impl SegmentTracker {
32 pub fn new() -> Self {
34 Self {
35 inner: Mutex::new(TrackerInner {
36 ref_counts: HashMap::new(),
37 pending_deletions: HashMap::new(),
38 scheduled_deletions: HashSet::new(),
39 }),
40 }
41 }
42
43 pub fn register(&self, segment_id: &str) {
45 let mut inner = self.inner.lock();
46 inner.ref_counts.entry(segment_id.to_string()).or_insert(0);
47 }
48
49 pub fn acquire(&self, segment_ids: &[String]) -> Vec<String> {
52 let mut inner = self.inner.lock();
53 let mut acquired = Vec::with_capacity(segment_ids.len());
54 for id in segment_ids {
55 if inner.pending_deletions.contains_key(id) || inner.scheduled_deletions.contains(id) {
56 continue;
57 }
58 *inner.ref_counts.entry(id.clone()).or_insert(0) += 1;
59 acquired.push(id.clone());
60 }
61 acquired
62 }
63
64 pub fn release(&self, segment_ids: &[String]) -> Vec<SegmentId> {
67 let mut inner = self.inner.lock();
68 let mut ready_for_deletion = Vec::new();
69
70 for id in segment_ids {
71 if let Some(count) = inner.ref_counts.get_mut(id) {
72 *count = count.saturating_sub(1);
73
74 if *count == 0
75 && let Some(segment_id) = inner.pending_deletions.remove(id)
76 {
77 inner.ref_counts.remove(id);
78 inner.scheduled_deletions.insert(id.clone());
79 ready_for_deletion.push(segment_id);
80 }
81 }
82 }
83
84 ready_for_deletion
85 }
86
87 pub fn mark_for_deletion(&self, segment_ids: &[String]) -> Vec<SegmentId> {
91 let mut inner = self.inner.lock();
92 let mut ready_for_deletion = Vec::new();
93
94 for id_str in segment_ids {
95 let Some(segment_id) = SegmentId::from_hex(id_str) else {
96 continue;
97 };
98
99 let Some(&ref_count) = inner.ref_counts.get(id_str) else {
101 continue;
102 };
103
104 if ref_count == 0 {
105 inner.ref_counts.remove(id_str);
106 inner.scheduled_deletions.insert(id_str.clone());
107 ready_for_deletion.push(segment_id);
108 } else {
109 inner.pending_deletions.insert(id_str.clone(), segment_id);
110 }
111 }
112
113 ready_for_deletion
114 }
115
116 pub fn is_pending_deletion(&self, segment_id: &str) -> bool {
118 self.inner.lock().pending_deletions.contains_key(segment_id)
119 }
120
121 pub fn is_deletion_protected(&self, segment_id: &str) -> bool {
128 let inner = self.inner.lock();
129 inner.pending_deletions.contains_key(segment_id)
130 || inner.scheduled_deletions.contains(segment_id)
131 }
132
133 pub fn complete_deletion(&self, segment_ids: &[SegmentId]) {
138 let mut inner = self.inner.lock();
139 for segment_id in segment_ids {
140 inner.scheduled_deletions.remove(&segment_id.to_hex());
141 }
142 }
143
144 pub fn ref_count(&self, segment_id: &str) -> usize {
146 self.inner
147 .lock()
148 .ref_counts
149 .get(segment_id)
150 .copied()
151 .unwrap_or(0)
152 }
153}
154
155impl Default for SegmentTracker {
156 fn default() -> Self {
157 Self::new()
158 }
159}
160
161pub struct SegmentSnapshot {
166 tracker: Arc<SegmentTracker>,
167 segment_ids: Vec<String>,
168 trained_vectors: Option<Arc<TrainedVectorStructures>>,
172 delete_fn: Option<Arc<dyn Fn(Vec<SegmentId>) + Send + Sync>>,
174}
175
176impl SegmentSnapshot {
177 pub fn new(tracker: Arc<SegmentTracker>, segment_ids: Vec<String>) -> Self {
179 Self {
180 tracker,
181 segment_ids,
182 trained_vectors: None,
183 delete_fn: None,
184 }
185 }
186
187 pub fn with_delete_fn(
189 tracker: Arc<SegmentTracker>,
190 segment_ids: Vec<String>,
191 delete_fn: Arc<dyn Fn(Vec<SegmentId>) + Send + Sync>,
192 ) -> Self {
193 Self {
194 tracker,
195 segment_ids,
196 trained_vectors: None,
197 delete_fn: Some(delete_fn),
198 }
199 }
200
201 pub(crate) fn with_generation(
204 tracker: Arc<SegmentTracker>,
205 segment_ids: Vec<String>,
206 trained_vectors: Option<Arc<TrainedVectorStructures>>,
207 delete_fn: Arc<dyn Fn(Vec<SegmentId>) + Send + Sync>,
208 ) -> Self {
209 Self {
210 tracker,
211 segment_ids,
212 trained_vectors,
213 delete_fn: Some(delete_fn),
214 }
215 }
216
217 pub fn segment_ids(&self) -> &[String] {
219 &self.segment_ids
220 }
221
222 pub(crate) fn trained_vectors(&self) -> Option<Arc<TrainedVectorStructures>> {
224 self.trained_vectors.clone()
225 }
226
227 pub fn is_empty(&self) -> bool {
229 self.segment_ids.is_empty()
230 }
231
232 pub fn len(&self) -> usize {
234 self.segment_ids.len()
235 }
236}
237
238impl Drop for SegmentSnapshot {
239 fn drop(&mut self) {
240 let to_delete = self.tracker.release(&self.segment_ids);
241 if !to_delete.is_empty() {
242 if let Some(delete_fn) = &self.delete_fn {
243 log::info!(
244 "[segment_snapshot] dropping snapshot, deleting {} deferred segments",
245 to_delete.len()
246 );
247 delete_fn(to_delete);
248 } else {
249 log::warn!(
250 "[segment_snapshot] {} segments ready for deletion but no delete_fn provided",
251 to_delete.len()
252 );
253 }
254 }
255 }
256}
257
258#[cfg(test)]
259mod tests {
260 use super::*;
261
262 const SEG1: &str = "00000000000000000000000000000001";
263 const SEG2: &str = "00000000000000000000000000000002";
264
265 #[test]
266 fn test_tracker_register_and_acquire() {
267 let tracker = SegmentTracker::new();
268
269 tracker.register(SEG1);
270 tracker.register(SEG2);
271
272 let acquired = tracker.acquire(&[SEG1.to_string(), SEG2.to_string()]);
273 assert_eq!(acquired.len(), 2);
274
275 assert_eq!(tracker.ref_count(SEG1), 1);
276 assert_eq!(tracker.ref_count(SEG2), 1);
277 }
278
279 #[test]
280 fn test_tracker_release() {
281 let tracker = SegmentTracker::new();
282
283 tracker.register(SEG1);
284 tracker.acquire(&[SEG1.to_string()]);
285 tracker.acquire(&[SEG1.to_string()]);
286
287 assert_eq!(tracker.ref_count(SEG1), 2);
288
289 tracker.release(&[SEG1.to_string()]);
290 assert_eq!(tracker.ref_count(SEG1), 1);
291
292 tracker.release(&[SEG1.to_string()]);
293 assert_eq!(tracker.ref_count(SEG1), 0);
294 }
295
296 #[test]
297 fn test_tracker_mark_for_deletion_no_refs() {
298 let tracker = SegmentTracker::new();
299
300 tracker.register(SEG1);
301
302 let ready = tracker.mark_for_deletion(&[SEG1.to_string()]);
303 assert_eq!(ready.len(), 1);
304 assert!(!tracker.is_pending_deletion(SEG1));
305 assert!(tracker.is_deletion_protected(SEG1));
306
307 tracker.complete_deletion(&ready);
308 assert!(!tracker.is_deletion_protected(SEG1));
309 }
310
311 #[test]
312 fn test_tracker_mark_for_deletion_with_refs() {
313 let tracker = SegmentTracker::new();
314
315 tracker.register(SEG1);
316 tracker.acquire(&[SEG1.to_string()]);
317
318 let ready = tracker.mark_for_deletion(&[SEG1.to_string()]);
319 assert!(ready.is_empty());
320 assert!(tracker.is_pending_deletion(SEG1));
321
322 let deleted = tracker.release(&[SEG1.to_string()]);
323 assert_eq!(deleted.len(), 1);
324 assert!(!tracker.is_pending_deletion(SEG1));
325 assert!(tracker.is_deletion_protected(SEG1));
326
327 tracker.complete_deletion(&deleted);
328 assert!(!tracker.is_deletion_protected(SEG1));
329 }
330
331 #[test]
332 fn test_tracker_double_mark_for_deletion() {
333 let tracker = SegmentTracker::new();
334 tracker.register(SEG1);
335
336 let ready1 = tracker.mark_for_deletion(&[SEG1.to_string()]);
337 assert_eq!(ready1.len(), 1);
338
339 let ready2 = tracker.mark_for_deletion(&[SEG1.to_string()]);
341 assert!(ready2.is_empty());
342 }
343
344 #[test]
345 fn test_tracker_acquire_unregistered() {
346 let tracker = SegmentTracker::new();
347
348 let acquired = tracker.acquire(&[SEG1.to_string()]);
350 assert_eq!(acquired.len(), 1);
351 assert_eq!(tracker.ref_count(SEG1), 1);
352 }
353
354 #[test]
355 fn test_tracker_release_without_acquire() {
356 let tracker = SegmentTracker::new();
357 tracker.register(SEG1);
358
359 let deleted = tracker.release(&[SEG1.to_string()]);
361 assert!(deleted.is_empty());
362 assert_eq!(tracker.ref_count(SEG1), 0);
363 }
364
365 #[test]
366 fn test_snapshot_drop_triggers_deferred_delete() {
367 use std::sync::atomic::{AtomicUsize, Ordering};
368
369 let tracker = Arc::new(SegmentTracker::new());
370 tracker.register(SEG1);
371 tracker.register(SEG2);
372
373 let delete_count = Arc::new(AtomicUsize::new(0));
374 let dc = Arc::clone(&delete_count);
375 let delete_fn: Arc<dyn Fn(Vec<SegmentId>) + Send + Sync> = Arc::new(move |ids| {
376 dc.fetch_add(ids.len(), Ordering::SeqCst);
377 });
378
379 let acquired = tracker.acquire(&[SEG1.to_string(), SEG2.to_string()]);
381 let snapshot =
382 SegmentSnapshot::with_delete_fn(Arc::clone(&tracker), acquired, Arc::clone(&delete_fn));
383
384 let ready = tracker.mark_for_deletion(&[SEG1.to_string(), SEG2.to_string()]);
386 assert!(ready.is_empty());
387 assert!(tracker.is_pending_deletion(SEG1));
388 assert!(tracker.is_pending_deletion(SEG2));
389
390 drop(snapshot);
392 assert_eq!(delete_count.load(Ordering::SeqCst), 2);
393 assert!(!tracker.is_pending_deletion(SEG1));
394 assert!(!tracker.is_pending_deletion(SEG2));
395 assert!(tracker.is_deletion_protected(SEG1));
396 assert!(tracker.is_deletion_protected(SEG2));
397
398 tracker.complete_deletion(&[
399 SegmentId::from_hex(SEG1).unwrap(),
400 SegmentId::from_hex(SEG2).unwrap(),
401 ]);
402 assert!(!tracker.is_deletion_protected(SEG1));
403 assert!(!tracker.is_deletion_protected(SEG2));
404 }
405}