Skip to main content

engramdb_io/
batch.rs

1//! 批式 badge 读取:排序去重 + 预取计划 + 多线程 gather。
2
3use std::collections::HashMap;
4use std::fs::File;
5use std::os::unix::fs::FileExt;
6use std::path::Path;
7
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
44/// 读路径(pread 池化):一次 `gather_plan` = 按 plan 拉取全部 badge 并组装 `[n,width]`。
45pub struct BadgeGather<'a> {
46    pub layout: &'a Layout,
47    files: Vec<File>,
48}
49
50impl<'a> BadgeGather<'a> {
51    pub fn open(dir: &Path, layout: &'a Layout) -> std::io::Result<Self> {
52        let n = layout.shards as usize;
53        let mut files = Vec::with_capacity(n);
54        for i in 0..n {
55            files.push(File::open(dir.join(format!("shard_{:03}.bin", i)))?);
56        }
57        Ok(Self { layout, files })
58    }
59
60    pub fn into_files(self) -> Vec<File> {
61        self.files
62    }
63
64    /// 多线程 gather:`keys`/`out` 按 chunk 分片并行(各线程独立 badge 缓冲区)。
65    pub fn gather_parallel(
66        &self,
67        keys: &[u64],
68        out: &mut [u8],
69        threads: usize,
70    ) -> std::io::Result<()> {
71        if threads <= 1 || keys.len() <= 1024 {
72            return self.gather_naive(keys, out);
73        }
74        let w = self.layout.width as usize;
75        let chunk_keys = keys.len().div_ceil(threads);
76        std::thread::scope(|s| {
77            for (kc, oc) in keys.chunks(chunk_keys).zip(out.chunks_mut(chunk_keys * w)) {
78                s.spawn(|| {
79                    let _ = self.gather_naive(kc, oc);
80                });
81            }
82        });
83        Ok(())
84    }
85
86    /// 朴素单线程:逐 key 定位 badge 读取(每个线程独立缓冲区)。
87    pub fn gather_naive(&self, keys: &[u64], out: &mut [u8]) -> std::io::Result<()> {
88        let w = self.layout.width as usize;
89        let rb = self.layout.row_bytes as usize;
90        let mut badge_buf = vec![0u8; self.layout.badge_bytes() as usize];
91        let mut last: Option<(u64, u64)> = None;
92        for (i, &k) in keys.iter().enumerate() {
93            let (shard, badge, in_badge) = self.layout.locate(k);
94            if last != Some((shard, badge)) {
95                let off = badge * self.layout.badge_bytes();
96                self.files[shard as usize].read_exact_at(&mut badge_buf, off)?;
97                last = Some((shard, badge));
98            }
99            let src = in_badge as usize * rb;
100            out[i * w..(i + 1) * w].copy_from_slice(&badge_buf[src..src + w]);
101        }
102        Ok(())
103    }
104
105    /// 页对齐专用读路径:按 shard 分线程,shard 内按键升序聚页(4KiB 对齐,每页只读一次),
106    /// 行跨页边界时补读该行尾部。各线程只写自己的采集缓冲,主线程回填 out(无共享 &mut)。
107    pub fn gather_pp(&self, keys: &[u64], out: &mut [u8], threads: usize) -> std::io::Result<()> {
108        const PAGE: u64 = 4096;
109        let w = self.layout.width as usize;
110        let rb = self.layout.row_bytes as usize;
111        let mut groups: HashMap<u64, Vec<(u64, usize)>> = HashMap::new();
112        for (i, &k) in keys.iter().enumerate() {
113            let (shard, _, _) = self.layout.locate(k);
114            groups.entry(shard).or_default().push((k, i));
115        }
116        let mut tasks: Vec<(u64, Vec<(u64, usize)>)> = groups.into_iter().collect();
117        tasks.sort_unstable_by_key(|&(s, _)| s);
118
119        // 各任务独立产出 (idxs 升序, rows 扁平)
120        let nt = threads.max(1).min(tasks.len());
121        let chunk = tasks.len().div_ceil(nt);
122        let mut results: Vec<(Vec<usize>, Vec<u8>)> = Vec::new();
123
124        std::thread::scope(|s| {
125            let mut handles = Vec::new();
126            let mut task_iter = tasks.into_iter();
127            while task_iter.len() > 0 {
128                let t: Vec<(u64, Vec<(u64, usize)>)> = task_iter.by_ref().take(chunk).collect();
129                handles.push(s.spawn(move || {
130                    let mut out_rows: Vec<u8> = Vec::new();
131                    let mut out_idxs: Vec<usize> = Vec::new();
132                    for (shard, mut pairs) in t {
133                        pairs.sort_unstable();
134                        let f = &self.files[shard as usize];
135                        let mut last_page: Option<u64> = None;
136                        let mut page = vec![0u8; (PAGE + 2 * (rb as u64)) as usize];
137                        let mut prev_key: Option<u64> = None;
138                        for (k, oi) in pairs {
139                            let (_, _, in_b) = self.layout.locate(k);
140                            let byte_off = k * rb as u64;
141                            let page_id = byte_off & !(PAGE - 1);
142                            if last_page != Some(page_id) {
143                                let want = (PAGE + rb as u64) as usize;
144                                let n = f.read_at(&mut page[..want], page_id).unwrap_or(0);
145                                let _ = n;
146                                last_page = Some(page_id);
147                            }
148                            let in_page = (byte_off - page_id) as usize;
149                            if in_page + rb <= PAGE as usize {
150                                out_rows.extend_from_slice(&page[in_page..in_page + rb]);
151                            } else {
152                                let mut tmp = vec![0u8; rb];
153                                let _ = f.read_exact_at(&mut tmp, byte_off);
154                                out_rows.extend_from_slice(&tmp);
155                            }
156                            out_idxs.push(oi);
157                            let _ = (in_b, prev_key);
158                            prev_key = Some(k);
159                        }
160                    }
161                    (out_idxs, out_rows)
162                }));
163            }
164            for h in handles {
165                if let Ok(r) = h.join() {
166                    results.push(r);
167                }
168            }
169        });
170
171        for (idxs, rows) in results {
172            for (j, &oi) in idxs.iter().enumerate() {
173                let slice = &rows[j * w..(j + 1) * w];
174                out[oi * w..(oi + 1) * w].copy_from_slice(slice);
175            }
176        }
177        Ok(())
178    }
179
180    /// 有序批式读:对同一 badge 的 keys 组内合并读(按 plan 的排序),
181    /// 期望预取服务器将 plan 先落地 —— 本实现直接做"计划->读取"。
182    pub fn gather_planned(&self, keys: &[u64], out: &mut [u8]) -> std::io::Result<()> {
183        // 简单实现:按 (shard,badge) 分组,逐组读一次,再按 key 顺序回填
184        let plan = PrefetchPlan::build(keys, self.layout);
185        let rb = self.layout.row_bytes as usize;
186        let w = self.layout.width as usize;
187        let mut cache: HashMap<(u64, u64), Vec<u8>> = HashMap::new();
188        let mut groups: HashMap<(u64, u64), Vec<usize>> = HashMap::new();
189        for (i, &k) in keys.iter().enumerate() {
190            let (s, b, _) = self.layout.locate(k);
191            groups.entry((s, b)).or_default().push(i);
192        }
193        for (&(s, b), idxs) in &groups {
194            let buf = cache.entry((s, b)).or_insert_with(|| {
195                let mut buf = vec![0u8; self.layout.badge_bytes() as usize];
196                let off = b * self.layout.badge_bytes();
197                let _ = self.files[s as usize].read_exact_at(&mut buf, off);
198                buf
199            });
200            for &i in idxs {
201                let (_, _, in_badge) = self.layout.locate(keys[i]);
202                let src = in_badge as usize * rb;
203                out[i * w..(i + 1) * w].copy_from_slice(&buf[src..src + w]);
204            }
205        }
206        let _ = plan;
207        Ok(())
208    }
209}
210
211#[cfg(test)]
212mod tests {
213    use super::*;
214    use engramdb_core::layout::Layout;
215
216    #[test]
217    fn plan_dedup_sort() {
218        let layout = Layout::new(1, 10_000, 160, 1); // 单分片
219        let keys = vec![9999, 0, 5, 9999, 250, 250];
220        let p = PrefetchPlan::build(&keys, &layout);
221        assert_eq!(p.badges(0), &[0, 10, 399]);
222    }
223}