1use std::collections::{HashMap, HashSet};
2use zer_core::{record::RecordId, traits::BlockIndex};
3
4pub struct InvertedIndex {
6 buckets: HashMap<String, Vec<RecordId>>,
7 record_keys: HashMap<RecordId, Vec<String>>,
8}
9
10impl InvertedIndex {
11 pub fn new() -> Self {
12 Self {
13 buckets: HashMap::new(),
14 record_keys: HashMap::new(),
15 }
16 }
17
18 pub fn insert(&mut self, record_id: RecordId, keys: Vec<String>) {
19 for key in &keys {
20 self.buckets.entry(key.clone()).or_default().push(record_id);
21 }
22 self.record_keys.insert(record_id, keys);
23 }
24
25 pub fn lookup_union(&self, keys: &[String], exclude: RecordId) -> Vec<RecordId> {
28 let mut seen: HashSet<RecordId> = HashSet::new();
29 for key in keys {
30 if let Some(ids) = self.buckets.get(key) {
31 for &id in ids {
32 if id != exclude {
33 seen.insert(id);
34 }
35 }
36 }
37 }
38 seen.into_iter().collect()
39 }
40
41 pub fn lookup_union_capped(
47 &self,
48 keys: &[String],
49 exclude: RecordId,
50 max_bucket_size: usize,
51 ) -> Vec<RecordId> {
52 let mut seen: HashSet<RecordId> = HashSet::new();
53 for key in keys {
54 if let Some(ids) = self.buckets.get(key) {
55 if max_bucket_size > 0 && ids.len() > max_bucket_size {
56 continue;
57 }
58 for &id in ids {
59 if id != exclude {
60 seen.insert(id);
61 }
62 }
63 }
64 }
65 seen.into_iter().collect()
66 }
67
68 pub fn bucket_size(&self, key: &str) -> usize {
70 self.buckets.get(key).map_or(0, |v| v.len())
71 }
72
73 pub fn oversized_buckets(&self, max_size: usize) -> usize {
75 self.buckets.values().filter(|v| v.len() > max_size).count()
76 }
77
78 pub fn all_pairs(
82 &self,
83 id_to_idx: &HashMap<RecordId, usize>,
84 max_bucket_size: usize,
85 ) -> Vec<(usize, usize)> {
86 let mut pairs: Vec<(usize, usize)> = Vec::new();
87 for bucket in self.buckets.values() {
88 if max_bucket_size > 0 && bucket.len() > max_bucket_size {
89 continue;
90 }
91 let indices: Vec<usize> = bucket
92 .iter()
93 .filter_map(|id| id_to_idx.get(id).copied())
94 .collect();
95 for a in 0..indices.len() {
96 for b in (a + 1)..indices.len() {
97 let (i, j) = (indices[a], indices[b]);
98 pairs.push(if i < j { (i, j) } else { (j, i) });
99 }
100 }
101 }
102 pairs.sort_unstable();
103 pairs.dedup();
104 pairs
105 }
106
107 pub fn remove(&mut self, record_id: RecordId) {
108 if let Some(keys) = self.record_keys.remove(&record_id) {
109 for key in keys {
110 if let Some(bucket) = self.buckets.get_mut(&key) {
111 bucket.retain(|&id| id != record_id);
112 }
113 }
114 }
115 }
116
117 pub fn len(&self) -> usize {
118 self.buckets.len()
119 }
120
121 pub fn is_empty(&self) -> bool {
122 self.buckets.is_empty()
123 }
124
125 pub fn record_count(&self) -> usize {
126 self.record_keys.len()
127 }
128}
129
130impl Default for InvertedIndex {
131 fn default() -> Self {
132 Self::new()
133 }
134}
135
136impl BlockIndex for InvertedIndex {
137 fn insert(&mut self, record_id: RecordId, keys: Vec<String>) {
138 self.insert(record_id, keys);
139 }
140
141 fn lookup_union(&self, keys: &[String], exclude: RecordId) -> Vec<RecordId> {
142 self.lookup_union(keys, exclude)
143 }
144
145 fn remove(&mut self, record_id: RecordId) {
146 self.remove(record_id);
147 }
148
149 fn as_any(&self) -> &dyn std::any::Any {
150 self
151 }
152
153 fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
154 self
155 }
156}
157
158#[cfg(test)]
159mod tests {
160 use super::*;
161
162 fn make_index() -> InvertedIndex {
163 let mut idx = InvertedIndex::new();
164 idx.insert(1, vec!["key_a".into(), "key_b".into()]);
165 idx.insert(2, vec!["key_b".into(), "key_c".into()]);
166 idx.insert(3, vec!["key_c".into(), "key_d".into()]);
167 idx
168 }
169
170 #[test]
171 fn lookup_union_returns_all_matching() {
172 let idx = make_index();
173 let mut result = idx.lookup_union(&["key_b".into()], 99);
174 result.sort();
175 assert_eq!(result, vec![1, 2]);
176 }
177
178 #[test]
179 fn lookup_union_deduplicates() {
180 let mut idx = InvertedIndex::new();
181 idx.insert(1, vec!["k1".into(), "k2".into()]);
182 idx.insert(2, vec!["k1".into(), "k2".into()]);
183
184 let result = idx.lookup_union(&["k1".into(), "k2".into()], 99);
185 assert_eq!(result.len(), 2);
186 }
187
188 #[test]
189 fn no_self_candidates() {
190 let idx = make_index();
191 let result = idx.lookup_union(&["key_a".into(), "key_b".into()], 1);
192 assert!(!result.contains(&1));
193 }
194
195 #[test]
196 fn remove_cleans_up() {
197 let mut idx = make_index();
198 idx.remove(1);
199 let result = idx.lookup_union(&["key_a".into(), "key_b".into()], 99);
200 assert!(!result.contains(&1));
201 }
202
203 #[test]
204 fn block_index_trait_insert_and_lookup() {
205 let mut idx: Box<dyn BlockIndex> = Box::new(InvertedIndex::new());
206 idx.insert(10, vec!["k".into()]);
207 idx.insert(20, vec!["k".into()]);
208 let mut result = idx.lookup_union(&["k".into()], 99);
209 result.sort();
210 assert_eq!(result, vec![10, 20]);
211 }
212
213 #[test]
214 fn block_index_trait_remove() {
215 let mut idx: Box<dyn BlockIndex> = Box::new(InvertedIndex::new());
216 idx.insert(1, vec!["x".into()]);
217 idx.remove(1);
218 let result = idx.lookup_union(&["x".into()], 99);
219 assert!(result.is_empty());
220 }
221
222 #[test]
223 fn lookup_union_capped_skips_oversized_bucket() {
224 let mut idx = InvertedIndex::new();
225 for id in 1u64..=5 {
227 idx.insert(id, vec!["big_key".into()]);
228 }
229 idx.insert(10u64, vec!["small_key".into()]);
230 idx.insert(11u64, vec!["small_key".into()]);
231
232 let result = idx.lookup_union_capped(&["big_key".into(), "small_key".into()], 99, 3);
234 assert!(!result.contains(&1), "big_key bucket must be skipped");
235 assert!(result.contains(&10), "small_key bucket must be included");
236 assert!(result.contains(&11), "small_key bucket must be included");
237 }
238
239 #[test]
240 fn lookup_union_capped_zero_cap_disables_limit() {
241 let mut idx = InvertedIndex::new();
242 for id in 1u64..=10 {
243 idx.insert(id, vec!["k".into()]);
244 }
245 let result = idx.lookup_union_capped(&["k".into()], 1, 0);
247 assert_eq!(result.len(), 9);
248 }
249
250 #[test]
251 fn oversized_buckets_count_is_correct() {
252 let mut idx = InvertedIndex::new();
253 for id in 1u64..=5 {
254 idx.insert(id, vec!["big".into()]);
255 }
256 idx.insert(10u64, vec!["small".into()]);
257 assert_eq!(idx.oversized_buckets(4), 1);
258 assert_eq!(idx.oversized_buckets(5), 0);
259 assert_eq!(idx.oversized_buckets(0), 2);
260 }
261}