1use std::collections::HashMap;
4use std::fs::File;
5use std::path::Path;
6
7use crate::backend::{default_backend, IoBackend};
8use engramdb_core::layout::Layout;
9
10#[derive(Debug, Default, Clone)]
12pub struct PrefetchPlan {
13 pub shard_badges: HashMap<u64, Vec<u64>>,
15 pub n_badges: usize,
16}
17
18impl PrefetchPlan {
19 pub fn build(keys: &[u64], layout: &Layout) -> Self {
20 let mut shard_badges: HashMap<u64, Vec<u64>> = HashMap::new();
21 for &k in keys {
22 let (shard, badge, _) = layout.locate(k);
23 shard_badges.entry(shard).or_default().push(badge);
24 }
25 for v in shard_badges.values_mut() {
26 v.sort_unstable();
27 v.dedup();
28 }
29 let n_badges = shard_badges.values().map(|v| v.len()).sum();
30 Self {
31 shard_badges,
32 n_badges,
33 }
34 }
35
36 pub fn badges(&self, shard: u64) -> &[u64] {
37 self.shard_badges
38 .get(&shard)
39 .map(|v| v.as_slice())
40 .unwrap_or(&[])
41 }
42
43 pub fn entry(&mut self, shard: u64, badge: u64) {
45 self.shard_badges.entry(shard).or_default().push(badge);
46 self.n_badges += 1;
47 }
48
49 pub fn settle(&mut self) {
51 let mut n = 0usize;
52 for v in self.shard_badges.values_mut() {
53 v.sort_unstable();
54 v.dedup();
55 n += v.len();
56 }
57 self.n_badges = n;
58 }
59}
60
61pub struct BadgeGather<'a> {
63 pub layout: &'a Layout,
64 files: Vec<File>,
65 backend: Box<dyn IoBackend>,
66}
67
68impl<'a> BadgeGather<'a> {
69 pub fn open(dir: &Path, layout: &'a Layout) -> std::io::Result<Self> {
70 Self::open_with_backend(dir, layout, default_backend())
71 }
72
73 pub fn open_with_backend(
75 dir: &Path,
76 layout: &'a Layout,
77 backend: Box<dyn IoBackend>,
78 ) -> std::io::Result<Self> {
79 let n = layout.shards as usize;
80 let mut files = Vec::with_capacity(n);
81 for i in 0..n {
82 let s = dir.join(format!("shard_{:03}.bin", i));
83 let b = dir.join(format!("badge_{:03}.bin", i));
84 let p = if s.exists() { s } else { b };
85 files.push(File::open(p)?);
86 }
87 Ok(Self {
88 layout,
89 files,
90 backend,
91 })
92 }
93
94 pub fn into_files(self) -> Vec<File> {
95 self.files
96 }
97
98 pub fn gather_parallel(
100 &self,
101 keys: &[u64],
102 out: &mut [u8],
103 threads: usize,
104 ) -> std::io::Result<()> {
105 if threads <= 1 || keys.len() <= 1024 {
106 return self.gather_naive(keys, out);
107 }
108 let w = self.layout.width as usize;
109 let chunk_keys = keys.len().div_ceil(threads);
110 std::thread::scope(|s| {
111 for (kc, oc) in keys.chunks(chunk_keys).zip(out.chunks_mut(chunk_keys * w)) {
112 s.spawn(|| {
113 let _ = self.gather_naive(kc, oc);
114 });
115 }
116 });
117 Ok(())
118 }
119
120 pub fn gather_naive(&self, keys: &[u64], out: &mut [u8]) -> std::io::Result<()> {
122 let w = self.layout.width as usize;
123 let rb = self.layout.row_bytes as usize;
124 let mut badge_buf = vec![0u8; self.layout.badge_bytes() as usize];
125 let mut last: Option<(u64, u64)> = None;
126 for (i, &k) in keys.iter().enumerate() {
127 let (shard, badge, in_badge) = self.layout.locate(k);
128 if last != Some((shard, badge)) {
129 let off = badge * self.layout.badge_bytes();
130 self.backend
131 .read_exact_at(&self.files[shard as usize], &mut badge_buf, off)?;
132 last = Some((shard, badge));
133 }
134 let src = in_badge as usize * rb;
135 out[i * w..(i + 1) * w].copy_from_slice(&badge_buf[src..src + w]);
136 }
137 Ok(())
138 }
139
140 pub fn gather_pp(&self, keys: &[u64], out: &mut [u8], threads: usize) -> std::io::Result<()> {
143 const PAGE: u64 = 4096;
144 let w = self.layout.width as usize;
145 let rb = self.layout.row_bytes as usize;
146 let mut groups: HashMap<u64, Vec<(u64, usize)>> = HashMap::new();
147 for (i, &k) in keys.iter().enumerate() {
148 let (shard, _, _) = self.layout.locate(k);
149 groups.entry(shard).or_default().push((k, i));
150 }
151 let mut tasks: Vec<(u64, Vec<(u64, usize)>)> = groups.into_iter().collect();
152 tasks.sort_unstable_by_key(|&(s, _)| s);
153
154 let nt = threads.max(1).min(tasks.len());
156 let chunk = tasks.len().div_ceil(nt);
157 let mut results: Vec<(Vec<usize>, Vec<u8>)> = Vec::new();
158
159 std::thread::scope(|s| {
160 let mut handles = Vec::new();
161 let mut task_iter = tasks.into_iter();
162 while task_iter.len() > 0 {
163 let t: Vec<(u64, Vec<(u64, usize)>)> = task_iter.by_ref().take(chunk).collect();
164 handles.push(s.spawn(move || {
165 let mut out_rows: Vec<u8> = Vec::new();
166 let mut out_idxs: Vec<usize> = Vec::new();
167 for (shard, mut pairs) in t {
168 pairs.sort_unstable();
169 let f = &self.files[shard as usize];
170 let mut last_page: Option<u64> = None;
171 let mut page = vec![0u8; (PAGE + 2 * (rb as u64)) as usize];
172 let mut prev_key: Option<u64> = None;
173 for (k, oi) in pairs {
174 let (_, _, in_b) = self.layout.locate(k);
175 let local_row = k % self.layout.rows_per_shard;
179 let byte_off = local_row * rb as u64;
180 let page_id = byte_off & !(PAGE - 1);
181 if last_page != Some(page_id) {
182 let want = (PAGE + rb as u64) as usize;
183 let n = self
184 .backend
185 .read_at(f, &mut page[..want], page_id)
186 .unwrap_or(0);
187 let _ = n;
188 last_page = Some(page_id);
189 }
190 let in_page = (byte_off - page_id) as usize;
191 if in_page + rb <= PAGE as usize {
192 out_rows.extend_from_slice(&page[in_page..in_page + rb]);
193 } else {
194 let mut tmp = vec![0u8; rb];
195 let _ = self.backend.read_exact_at(f, &mut tmp, byte_off);
196 out_rows.extend_from_slice(&tmp);
197 }
198 out_idxs.push(oi);
199 let _ = (in_b, prev_key);
200 prev_key = Some(k);
201 }
202 }
203 (out_idxs, out_rows)
204 }));
205 }
206 for h in handles {
207 if let Ok(r) = h.join() {
208 results.push(r);
209 }
210 }
211 });
212
213 for (idxs, rows) in results {
214 for (j, &oi) in idxs.iter().enumerate() {
215 let slice = &rows[j * w..(j + 1) * w];
216 out[oi * w..(oi + 1) * w].copy_from_slice(slice);
217 }
218 }
219 Ok(())
220 }
221
222 pub fn gather_planned(&self, keys: &[u64], out: &mut [u8]) -> std::io::Result<()> {
225 let plan = PrefetchPlan::build(keys, self.layout);
227 let rb = self.layout.row_bytes as usize;
228 let w = self.layout.width as usize;
229 let mut cache: HashMap<(u64, u64), Vec<u8>> = HashMap::new();
230 let mut groups: HashMap<(u64, u64), Vec<usize>> = HashMap::new();
231 for (i, &k) in keys.iter().enumerate() {
232 let (s, b, _) = self.layout.locate(k);
233 groups.entry((s, b)).or_default().push(i);
234 }
235 for (&(s, b), idxs) in &groups {
236 let buf = cache.entry((s, b)).or_insert_with(|| {
237 let mut buf = vec![0u8; self.layout.badge_bytes() as usize];
238 let off = b * self.layout.badge_bytes();
239 let _ = self
240 .backend
241 .read_exact_at(&self.files[s as usize], &mut buf, off);
242 buf
243 });
244 for &i in idxs {
245 let (_, _, in_badge) = self.layout.locate(keys[i]);
246 let src = in_badge as usize * rb;
247 out[i * w..(i + 1) * w].copy_from_slice(&buf[src..src + w]);
248 }
249 }
250 let _ = plan;
251 Ok(())
252 }
253 pub fn gather_plan(
257 &self,
258 keys: &[u64],
259 plan: &PrefetchPlan,
260 out: &mut [u8],
261 ) -> std::io::Result<()> {
262 let rb = self.layout.row_bytes as usize;
263 let w = self.layout.width as usize;
264 let mut groups: HashMap<(u64, u64), Vec<usize>> = HashMap::new();
265 for (i, &k) in keys.iter().enumerate() {
266 let (s, b, _) = self.layout.locate(k);
267 if plan.badges(s).contains(&b) {
268 groups.entry((s, b)).or_default().push(i);
269 } else {
270 let mut tmp = vec![0u8; rb];
272 self.backend
273 .read_exact_at(&self.files[s as usize], &mut tmp, k * rb as u64)?;
274 out[i * w..(i + 1) * w].copy_from_slice(&tmp);
275 }
276 }
277 for (&(s, b), idxs) in &groups {
278 let mut buf = vec![0u8; self.layout.badge_bytes() as usize];
279 let off = b * self.layout.badge_bytes();
280 self.backend
281 .read_exact_at(&self.files[s as usize], &mut buf, off)?;
282 for &i in idxs {
283 let (_, _, in_badge) = self.layout.locate(keys[i]);
284 let src = in_badge as usize * rb;
285 out[i * w..(i + 1) * w].copy_from_slice(&buf[src..src + w]);
286 }
287 }
288 Ok(())
289 }
290}
291
292#[cfg(test)]
293mod tests {
294 use super::*;
295 use engramdb_core::layout::Layout;
296
297 #[test]
298 fn plan_dedup_sort() {
299 let layout = Layout::new(1, 10_000, 160, 1); let keys = vec![9999, 0, 5, 9999, 250, 250];
301 let p = PrefetchPlan::build(&keys, &layout);
302 assert_eq!(p.badges(0), &[0, 10, 399]);
303 }
304
305 #[test]
306 fn gather_plan_end_to_end() {
307 use std::io::Write;
308 let dir = std::env::temp_dir().join("engramdb-gather-plan-test");
309 let _ = std::fs::remove_dir_all(&dir);
310 std::fs::create_dir_all(&dir).unwrap();
311 let layout = Layout::new(1, 100, 8, 1);
313 let mut f = std::fs::File::create(dir.join("shard_000.bin")).unwrap();
314 for i in 0..512u64 {
315 f.write_all(&i.to_le_bytes()).unwrap();
316 }
317 drop(f);
318 let bg = BadgeGather::open(&dir, &layout).unwrap();
319 let keys = vec![3u64, 77, 3, 99, 60];
320 let plan = PrefetchPlan::build(&keys, &layout);
321 let mut out = vec![0u8; keys.len() * 8];
322 bg.gather_plan(&keys, &plan, &mut out).unwrap();
323 for (j, &want) in keys.iter().enumerate() {
324 let got = u64::from_le_bytes(out[j * 8..(j + 1) * 8].try_into().unwrap());
325 assert_eq!(got, want, "rowid {want} at out[{j}]");
326 }
327 let _ = std::fs::remove_dir_all(&dir);
328 }
329
330 #[test]
331 fn gather_pp_multishard_uses_local_row_offsets() {
332 use std::io::Write;
333 let dir = std::env::temp_dir().join("engramdb-gather-pp-multishard-test");
334 let _ = std::fs::remove_dir_all(&dir);
335 std::fs::create_dir_all(&dir).unwrap();
336 let shards = 3u64;
337 let rows_per_shard = 32u64;
338 let width = 4u64;
339 let layout = Layout::new(shards, rows_per_shard, width, 1);
340 for s in 0..shards {
341 let mut f = std::fs::File::create(dir.join(format!("shard_{s:03}.bin"))).unwrap();
342 for local in 0..rows_per_shard {
343 let val = (s * 1_000_000 + local) as u32;
344 f.write_all(&val.to_le_bytes()).unwrap();
345 }
346 }
347 let bg = BadgeGather::open(&dir, &layout).unwrap();
348 let keys = vec![0u64, 1, 31, 32, 63, 64, 95];
349 let mut out = vec![0u8; keys.len() * width as usize];
350 bg.gather_pp(&keys, &mut out, 4).unwrap();
351 for (j, &k) in keys.iter().enumerate() {
352 let shard = k / rows_per_shard;
353 let local = k % rows_per_shard;
354 let want = (shard * 1_000_000 + local) as u32;
355 let got = u32::from_le_bytes(out[j * 4..(j + 1) * 4].try_into().unwrap());
356 assert_eq!(got, want, "gather_pp rowid {k} at out[{j}]");
357 }
358 let _ = std::fs::remove_dir_all(&dir);
359 }
360}