Skip to main content

engramdb_io/
batch.rs

1//! 批式 badge 读取:排序去重 + 预取计划 + 多线程 gather。
2
3use 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/// 预取计划:按分片分组的 badge 块列表(每 shard 内部升序,供顺序预读/合并)。
11#[derive(Debug, Default, Clone)]
12pub struct PrefetchPlan {
13    /// shard_id -> 去重升序的 badge 块号
14    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    /// 追加一个(shard, badge)到计划(流水线增量累积;重复由 settle 处理)。
44    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    /// 结算:shard 内排序 + 去重 + 重算 n_badges(计划被消费前调用)。
50    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
61/// 读路径(pread 池化):一次 `gather_plan` = 按 plan 拉取全部 badge 并组装 `[n,width]`。
62pub 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    /// 可注入后端(测试/基准:Preadv;未来 Linux:io_uring)。
74    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    /// 多线程 gather:`keys`/`out` 按 chunk 分片并行(各线程独立 badge 缓冲区)。
99    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    /// 朴素单线程:逐 key 定位 badge 读取(每个线程独立缓冲区)。
121    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    /// 页对齐专用读路径:按 shard 分线程,shard 内按键升序聚页(4KiB 对齐,每页只读一次),
141    /// 行跨页边界时补读该行尾部。各线程只写自己的采集缓冲,主线程回填 out(无共享 &mut)。
142    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        // 各任务独立产出 (idxs 升序, rows 扁平)
155        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                            // gather_pp groups by shard and reads from that
176                            // shard's file, so the byte offset must be local
177                            // to the shard rather than the global rowid.
178                            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    /// 有序批式读:对同一 badge 的 keys 组内合并读(按 plan 的排序),
223    /// 期望预取服务器将 plan 先落地 —— 本实现直接做"计划->读取"。
224    pub fn gather_planned(&self, keys: &[u64], out: &mut [u8]) -> std::io::Result<()> {
225        // 简单实现:按 (shard,badge) 分组,逐组读一次,再按 key 顺序回填
226        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    /// 按计划消费:`plan` 决定需要读的 badge 集合,本函数按键组读取一次完整 badge
254    /// 并回填 `out`(与 `keys` 顺序一致)。与 `gather_planned` 的差异:计划预先生成
255    /// (由 StreamingPlanner/外部预取服务器),这里只做"计划 -> 数据"。
256    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                // 计划外的 key:退化为直接行读(防御;正常流式路径不会出现)
271                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); // 单分片
300        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        // 单分片 100 行 × 8 字节(u64 值 = rowid);文件填满完整 badge(512 行 × 8B = 4096B)
312        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}