1use std::sync::Arc;
13
14use parking_lot::RwLock;
15use rustc_hash::FxHashMap;
16
17use crate::dsl::Field;
18use crate::segment::SegmentReader;
19
20pub struct LazyGlobalStats {
26 segments: Vec<Arc<SegmentReader>>,
28 total_docs: u64,
30 sparse_idf_cache: RwLock<FxHashMap<u32, FxHashMap<u32, f32>>>,
32 sparse_total_vectors_cache: RwLock<FxHashMap<u32, u64>>,
34 text_idf_cache: RwLock<FxHashMap<u32, FxHashMap<String, f32>>>,
36 avg_field_len_cache: RwLock<FxHashMap<u32, f32>>,
38 text_df_cache: RwLock<FxHashMap<u32, FxHashMap<Vec<u8>, u64>>>,
41 text_corpus_cache: RwLock<FxHashMap<u32, u64>>,
44}
45
46impl LazyGlobalStats {
47 pub fn new(segments: Vec<Arc<SegmentReader>>) -> Self {
49 let total_docs: u64 = segments.iter().map(|s| s.num_docs() as u64).sum();
50 Self {
51 segments,
52 total_docs,
53 sparse_idf_cache: RwLock::new(FxHashMap::default()),
54 sparse_total_vectors_cache: RwLock::new(FxHashMap::default()),
55 text_idf_cache: RwLock::new(FxHashMap::default()),
56 avg_field_len_cache: RwLock::new(FxHashMap::default()),
57 text_df_cache: RwLock::new(FxHashMap::default()),
58 text_corpus_cache: RwLock::new(FxHashMap::default()),
59 }
60 }
61
62 pub fn text_df(&self, field: Field, term: &[u8]) -> u64 {
67 {
68 let cache = self.text_df_cache.read();
69 if let Some(field_cache) = cache.get(&field.0)
70 && let Some(&df) = field_cache.get(term)
71 {
72 return df;
73 }
74 }
75 let df = self.compute_text_df_bytes(field, term);
76 self.text_df_cache
77 .write()
78 .entry(field.0)
79 .or_default()
80 .insert(term.to_vec(), df);
81 df
82 }
83
84 pub fn text_corpus_size(&self, field: Field) -> u64 {
87 if let Some(&size) = self.text_corpus_cache.read().get(&field.0) {
88 return size;
89 }
90 let size: u64 = self
91 .segments
92 .iter()
93 .map(|segment| segment.text_corpus_size(field) as u64)
94 .sum();
95 self.text_corpus_cache.write().insert(field.0, size);
96 size
97 }
98
99 pub fn text_stats_for(&self, terms: &[(Field, Vec<u8>)]) -> GlobalStats {
104 let mut builder = GlobalStatsBuilder::new();
105 builder.total_docs = self.total_docs;
106 let mut fields_seen: FxHashMap<u32, ()> = FxHashMap::default();
107 for (field, term) in terms {
108 if fields_seen.insert(field.0, ()).is_none() {
109 builder.set_avg_field_len(*field, self.avg_field_len(*field));
110 builder.set_text_corpus_size(*field, self.text_corpus_size(*field));
111 }
112 let df = self.text_df(*field, term);
113 if df > 0 {
114 builder.add_text_df(*field, String::from_utf8_lossy(term).into_owned(), df);
115 }
116 }
117 builder.build(0)
118 }
119
120 #[cfg(feature = "sync")]
121 fn compute_text_df_bytes(&self, field: Field, term: &[u8]) -> u64 {
122 self.segments
123 .iter()
124 .map(|segment| segment.text_doc_freq_sync(field, term).unwrap_or(0) as u64)
125 .sum()
126 }
127
128 #[cfg(not(feature = "sync"))]
129 fn compute_text_df_bytes(&self, _field: Field, _term: &[u8]) -> u64 {
130 0
131 }
132
133 #[inline]
135 pub fn total_docs(&self) -> u64 {
136 self.total_docs
137 }
138
139 pub fn sparse_idf(&self, field: Field, dim_id: u32) -> f32 {
143 {
145 let cache = self.sparse_idf_cache.read();
146 if let Some(field_cache) = cache.get(&field.0)
147 && let Some(&idf) = field_cache.get(&dim_id)
148 {
149 return idf;
150 }
151 }
152
153 let df = self.compute_sparse_df(field, dim_id);
155 let n = self.cached_sparse_n(field);
156 let idf = if df > 0 && n > 0 {
157 (n as f32 / df as f32).ln().max(0.0)
158 } else {
159 0.0
160 };
161
162 {
164 let mut cache = self.sparse_idf_cache.write();
165 cache.entry(field.0).or_default().insert(dim_id, idf);
166 }
167
168 idf
169 }
170
171 pub fn sparse_idf_weights(&self, field: Field, dim_ids: &[u32]) -> Vec<f32> {
176 let mut result = vec![0.0f32; dim_ids.len()];
178 let mut misses: Vec<usize> = Vec::new();
179 {
180 let cache = self.sparse_idf_cache.read();
181 if let Some(field_cache) = cache.get(&field.0) {
182 for (i, &dim_id) in dim_ids.iter().enumerate() {
183 if let Some(&idf) = field_cache.get(&dim_id) {
184 result[i] = idf;
185 } else {
186 misses.push(i);
187 }
188 }
189 } else {
190 misses.extend(0..dim_ids.len());
191 }
192 }
193
194 if misses.is_empty() {
195 return result;
196 }
197
198 let n = self.cached_sparse_n(field);
200
201 let mut new_entries: Vec<(u32, f32)> = Vec::with_capacity(misses.len());
203 for &i in &misses {
204 let dim_id = dim_ids[i];
205 let df = self.compute_sparse_df(field, dim_id);
206 let idf = if df > 0 && n > 0 {
207 (n as f32 / df as f32).ln().max(0.0)
208 } else {
209 0.0
210 };
211 result[i] = idf;
212 new_entries.push((dim_id, idf));
213 }
214
215 {
217 let mut cache = self.sparse_idf_cache.write();
218 let field_cache = cache.entry(field.0).or_default();
219 for (dim_id, idf) in new_entries {
220 field_cache.insert(dim_id, idf);
221 }
222 }
223
224 result
225 }
226
227 fn cached_sparse_n(&self, field: Field) -> u64 {
230 {
232 let cache = self.sparse_total_vectors_cache.read();
233 if let Some(&tv) = cache.get(&field.0) {
234 return tv.max(self.total_docs);
235 }
236 }
237 let tv = self.compute_sparse_total_vectors(field);
239 self.sparse_total_vectors_cache.write().insert(field.0, tv);
240 tv.max(self.total_docs)
241 }
242
243 pub fn text_idf(&self, field: Field, term: &str) -> f32 {
247 {
249 let cache = self.text_idf_cache.read();
250 if let Some(field_cache) = cache.get(&field.0)
251 && let Some(&idf) = field_cache.get(term)
252 {
253 return idf;
254 }
255 }
256
257 let df = self.compute_text_df(field, term);
260 let n = self.text_corpus_size(field) as f32;
261 let df_f = df as f32;
262 let idf = if df > 0 {
263 ((n - df_f + 0.5) / (df_f + 0.5) + 1.0).ln()
264 } else {
265 0.0
266 };
267
268 {
270 let mut cache = self.text_idf_cache.write();
271 cache
272 .entry(field.0)
273 .or_default()
274 .insert(term.to_string(), idf);
275 }
276
277 idf
278 }
279
280 pub fn avg_field_len(&self, field: Field) -> f32 {
282 {
284 let cache = self.avg_field_len_cache.read();
285 if let Some(&avg) = cache.get(&field.0) {
286 return avg;
287 }
288 }
289
290 let mut weighted_sum = 0.0f64;
292 let mut total_weight = 0u64;
293
294 for segment in &self.segments {
295 let avg_len = segment.avg_field_len(field);
296 let doc_count = segment.text_corpus_size(field) as u64;
298 if avg_len > 0.0 && doc_count > 0 {
299 weighted_sum += avg_len as f64 * doc_count as f64;
300 total_weight += doc_count;
301 }
302 }
303
304 let avg = if total_weight > 0 {
305 (weighted_sum / total_weight as f64) as f32
306 } else {
307 1.0
308 };
309
310 {
312 let mut cache = self.avg_field_len_cache.write();
313 cache.insert(field.0, avg);
314 }
315
316 avg
317 }
318
319 fn compute_sparse_df(&self, field: Field, dim_id: u32) -> u64 {
322 let mut df = 0u64;
323 for segment in &self.segments {
324 if let Some(sparse_index) = segment.seismic_index(field) {
325 df += sparse_index.doc_count(dim_id) as u64;
326 } else if let Some(sparse_index) = segment.sparse_indexes().get(&field.0) {
327 df += sparse_index.doc_count(dim_id) as u64;
328 }
329 }
330 df
331 }
332
333 fn compute_sparse_total_vectors(&self, field: Field) -> u64 {
336 let mut total = 0u64;
337 for segment in &self.segments {
338 if let Some(sparse_index) = segment.seismic_index(field) {
339 total += sparse_index.total_vectors() as u64;
340 } else if let Some(sparse_index) = segment.sparse_indexes().get(&field.0) {
341 total += sparse_index.total_vectors as u64;
342 }
343 }
344 total
345 }
346
347 fn compute_text_df(&self, field: Field, term: &str) -> u64 {
349 self.text_df(field, term.as_bytes())
350 }
351
352 pub fn num_segments(&self) -> usize {
354 self.segments.len()
355 }
356}
357
358impl std::fmt::Debug for LazyGlobalStats {
359 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
360 f.debug_struct("LazyGlobalStats")
361 .field("total_docs", &self.total_docs)
362 .field("num_segments", &self.segments.len())
363 .field("sparse_cache_fields", &self.sparse_idf_cache.read().len())
364 .field("text_cache_fields", &self.text_idf_cache.read().len())
365 .finish()
366 }
367}
368
369#[derive(Debug)]
373pub struct GlobalStats {
374 total_docs: u64,
376 sparse_stats: FxHashMap<u32, SparseFieldStats>,
378 text_stats: FxHashMap<u32, TextFieldStats>,
380 generation: u64,
382}
383
384#[derive(Debug, Default)]
386pub struct SparseFieldStats {
387 pub doc_freqs: FxHashMap<u32, u64>,
389}
390
391#[derive(Debug, Default)]
393pub struct TextFieldStats {
394 pub doc_freqs: FxHashMap<String, u64>,
396 pub avg_field_len: f32,
398 pub corpus_size: u64,
401}
402
403impl GlobalStats {
404 pub fn new() -> Self {
406 Self {
407 total_docs: 0,
408 sparse_stats: FxHashMap::default(),
409 text_stats: FxHashMap::default(),
410 generation: 0,
411 }
412 }
413
414 #[inline]
416 pub fn total_docs(&self) -> u64 {
417 self.total_docs
418 }
419
420 #[inline]
422 pub fn sparse_idf(&self, field: Field, dim_id: u32) -> f32 {
423 if let Some(stats) = self.sparse_stats.get(&field.0)
424 && let Some(&df) = stats.doc_freqs.get(&dim_id)
425 && df > 0
426 {
427 return (self.total_docs as f32 / df as f32).ln();
428 }
429 0.0
430 }
431
432 pub fn sparse_idf_weights(&self, field: Field, dim_ids: &[u32]) -> Vec<f32> {
434 dim_ids.iter().map(|&d| self.sparse_idf(field, d)).collect()
435 }
436
437 #[inline]
439 pub fn text_idf(&self, field: Field, term: &str) -> f32 {
440 if let Some(stats) = self.text_stats.get(&field.0)
441 && let Some(&df) = stats.doc_freqs.get(term)
442 {
443 let n = if stats.corpus_size > 0 {
444 stats.corpus_size as f32
445 } else {
446 self.total_docs as f32
447 };
448 let df = df as f32;
449 return ((n - df + 0.5) / (df + 0.5) + 1.0).ln();
450 }
451 0.0
452 }
453
454 pub fn text_df(&self, field: Field, term: &str) -> Option<u64> {
456 self.text_stats
457 .get(&field.0)
458 .and_then(|stats| stats.doc_freqs.get(term).copied())
459 }
460
461 pub fn text_corpus_size(&self, field: Field) -> u64 {
463 self.text_stats
464 .get(&field.0)
465 .map_or(0, |stats| stats.corpus_size)
466 }
467
468 pub fn text_fields(&self) -> impl Iterator<Item = (Field, &TextFieldStats)> {
470 self.text_stats
471 .iter()
472 .map(|(id, stats)| (Field(*id), stats))
473 }
474
475 #[inline]
477 pub fn avg_field_len(&self, field: Field) -> f32 {
478 self.text_stats
479 .get(&field.0)
480 .map(|s| s.avg_field_len)
481 .unwrap_or(1.0)
482 }
483
484 #[inline]
486 pub fn generation(&self) -> u64 {
487 self.generation
488 }
489}
490
491impl Default for GlobalStats {
492 fn default() -> Self {
493 Self::new()
494 }
495}
496
497pub struct GlobalStatsBuilder {
499 pub total_docs: u64,
501 sparse_stats: FxHashMap<u32, SparseFieldStats>,
502 text_stats: FxHashMap<u32, TextFieldStats>,
503}
504
505impl GlobalStatsBuilder {
506 pub fn new() -> Self {
508 Self {
509 total_docs: 0,
510 sparse_stats: FxHashMap::default(),
511 text_stats: FxHashMap::default(),
512 }
513 }
514
515 pub fn add_segment(&mut self, reader: &SegmentReader) {
517 self.total_docs += reader.num_docs() as u64;
518
519 }
522
523 pub fn add_sparse_df(&mut self, field: Field, dim_id: u32, doc_count: u64) {
525 let stats = self.sparse_stats.entry(field.0).or_default();
526 *stats.doc_freqs.entry(dim_id).or_insert(0) += doc_count;
527 }
528
529 pub fn add_text_df(&mut self, field: Field, term: String, doc_count: u64) {
531 let stats = self.text_stats.entry(field.0).or_default();
532 *stats.doc_freqs.entry(term).or_insert(0) += doc_count;
533 }
534
535 pub fn set_avg_field_len(&mut self, field: Field, avg_len: f32) {
537 let stats = self.text_stats.entry(field.0).or_default();
538 stats.avg_field_len = avg_len;
539 }
540
541 pub fn set_text_corpus_size(&mut self, field: Field, corpus_size: u64) {
543 let stats = self.text_stats.entry(field.0).or_default();
544 stats.corpus_size = corpus_size;
545 }
546
547 pub fn build(self, generation: u64) -> GlobalStats {
549 GlobalStats {
550 total_docs: self.total_docs,
551 sparse_stats: self.sparse_stats,
552 text_stats: self.text_stats,
553 generation,
554 }
555 }
556}
557
558impl Default for GlobalStatsBuilder {
559 fn default() -> Self {
560 Self::new()
561 }
562}
563
564pub struct GlobalStatsCache {
569 stats: RwLock<Option<Arc<GlobalStats>>>,
571 generation: RwLock<u64>,
573}
574
575impl GlobalStatsCache {
576 pub fn new() -> Self {
578 Self {
579 stats: RwLock::new(None),
580 generation: RwLock::new(0),
581 }
582 }
583
584 pub fn invalidate(&self) {
586 let mut current_gen = self.generation.write();
587 *current_gen += 1;
588 let mut stats = self.stats.write();
589 *stats = None;
590 }
591
592 pub fn generation(&self) -> u64 {
594 *self.generation.read()
595 }
596
597 pub fn get(&self) -> Option<Arc<GlobalStats>> {
599 self.stats.read().clone()
600 }
601
602 pub fn set(&self, stats: GlobalStats) {
604 let mut cached = self.stats.write();
605 *cached = Some(Arc::new(stats));
606 }
607
608 pub fn get_or_compute<F>(&self, compute: F) -> Arc<GlobalStats>
612 where
613 F: FnOnce(&mut GlobalStatsBuilder),
614 {
615 if let Some(stats) = self.get() {
617 return stats;
618 }
619
620 let current_gen = self.generation();
622 let mut builder = GlobalStatsBuilder::new();
623 compute(&mut builder);
624 let stats = Arc::new(builder.build(current_gen));
625
626 let mut cached = self.stats.write();
628 *cached = Some(Arc::clone(&stats));
629
630 stats
631 }
632
633 pub fn needs_rebuild(&self) -> bool {
635 self.stats.read().is_none()
636 }
637
638 pub fn set_stats(&self, stats: GlobalStats) {
640 let mut cached = self.stats.write();
641 *cached = Some(Arc::new(stats));
642 }
643}
644
645impl Default for GlobalStatsCache {
646 fn default() -> Self {
647 Self::new()
648 }
649}
650
651#[cfg(test)]
652mod tests {
653 use super::*;
654
655 #[test]
656 fn test_sparse_idf_computation() {
657 let mut builder = GlobalStatsBuilder::new();
658 builder.total_docs = 1000;
659 builder.add_sparse_df(Field(0), 42, 100); builder.add_sparse_df(Field(0), 43, 10); let stats = builder.build(1);
663
664 let idf_42 = stats.sparse_idf(Field(0), 42);
666 let idf_43 = stats.sparse_idf(Field(0), 43);
667
668 assert!(idf_43 > idf_42);
670 assert!((idf_42 - (1000.0_f32 / 100.0).ln()).abs() < 0.001);
671 assert!((idf_43 - (1000.0_f32 / 10.0).ln()).abs() < 0.001);
672 }
673
674 #[test]
675 fn test_text_idf_computation() {
676 let mut builder = GlobalStatsBuilder::new();
677 builder.total_docs = 10000;
678 builder.add_text_df(Field(0), "common".to_string(), 5000);
679 builder.add_text_df(Field(0), "rare".to_string(), 10);
680
681 let stats = builder.build(1);
682
683 let idf_common = stats.text_idf(Field(0), "common");
684 let idf_rare = stats.text_idf(Field(0), "rare");
685
686 assert!(idf_rare > idf_common);
688 }
689
690 #[test]
691 fn test_cache_invalidation() {
692 let cache = GlobalStatsCache::new();
693
694 assert!(cache.get().is_none());
696
697 let stats = cache.get_or_compute(|builder| {
699 builder.total_docs = 100;
700 });
701 assert_eq!(stats.total_docs(), 100);
702
703 assert!(cache.get().is_some());
705
706 cache.invalidate();
708 assert!(cache.get().is_none());
709 }
710}