1use super::{ChunkRef, ChunkedPayload, Manifest};
8use crate::crypto::{CryptoProvider, DefaultProvider};
9use crate::{ChunkId, DdError, ErrorCode, HashAlgorithm, Result};
10use reed_solomon_erasure::galois_8::ReedSolomon;
11use serde::{Deserialize, Serialize};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
14pub struct ErasureSpec {
15 pub data_shards: u32,
16 pub parity_shards: u32,
17}
18
19#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
20pub struct ErasureInfo {
21 pub spec: ErasureSpec,
22 pub data_count: u32,
24}
25
26impl ErasureSpec {
27 pub fn validate(self) -> Result<()> {
28 if self.data_shards == 0 || self.data_shards > 64 {
29 return Err(DdError::protocol(
30 ErrorCode::Ddp1006LimitExceeded,
31 "erasure data_shards must be 1..=64",
32 ));
33 }
34 if self.parity_shards == 0 || self.parity_shards > 16 {
35 return Err(DdError::protocol(
36 ErrorCode::Ddp1006LimitExceeded,
37 "erasure parity_shards must be 1..=16",
38 ));
39 }
40 Ok(())
41 }
42}
43
44pub fn apply_erasure(mut body: ChunkedPayload, spec: ErasureSpec) -> Result<ChunkedPayload> {
46 spec.validate()?;
47 let data_count = body.chunks.len() as u32;
48 if data_count == 0 {
49 return Ok(body);
50 }
51 let p = DefaultProvider;
52 let data_n = spec.data_shards as usize;
53 let parity_n = spec.parity_shards as usize;
54 let data_chunks = body.chunks.clone();
55 let mut i = 0usize;
56 while i < data_chunks.len() {
57 let end = (i + data_n).min(data_chunks.len());
58 let group = &data_chunks[i..end];
59 let width = group.iter().map(|c| c.len()).max().unwrap_or(0);
60 let mut shards = vec![vec![0u8; width]; data_n + parity_n];
61 for (s, chunk) in shards.iter_mut().zip(group.iter()) {
62 s[..chunk.len()].copy_from_slice(chunk);
63 }
64 if parity_n == 1 {
65 xor_parity(&mut shards, data_n);
66 } else {
67 let rs = ReedSolomon::new(data_n, parity_n).map_err(|e| {
68 DdError::protocol(
69 ErrorCode::Ddp1006LimitExceeded,
70 format!("reed-solomon: {e}"),
71 )
72 })?;
73 rs.encode(&mut shards)
74 .map_err(|e| DdError::protocol(ErrorCode::Dds2002CorruptChunk, format!("{e}")))?;
75 }
76 for shard in shards.iter().skip(data_n) {
77 let digest = p.hash(HashAlgorithm::Blake3, shard);
78 body.manifest.chunks.push(ChunkRef {
79 id: ChunkId::blake3(digest.0),
80 length: shard.len() as u32,
81 });
82 body.chunks.push(shard.clone());
83 }
84 i = end;
85 }
86 body.manifest.erasure = Some(ErasureInfo { spec, data_count });
87 Ok(body)
88}
89
90fn xor_parity(shards: &mut [Vec<u8>], data_n: usize) {
91 let width = shards.first().map(|s| s.len()).unwrap_or(0);
92 let mut parity = vec![0u8; width];
93 for shard in shards.iter().take(data_n) {
94 for (p, b) in parity.iter_mut().zip(shard.iter()) {
95 *p ^= *b;
96 }
97 }
98 if let Some(slot) = shards.get_mut(data_n) {
99 *slot = parity;
100 }
101}
102
103pub fn can_recover(manifest: &Manifest, present: &[bool]) -> bool {
104 let Some(info) = &manifest.erasure else {
105 return !present.is_empty() && present.iter().all(|x| *x);
106 };
107 let data_n = info.spec.data_shards as usize;
108 let parity_n = info.spec.parity_shards as usize;
109 let data_count = info.data_count as usize;
110 let stripe = data_n + parity_n;
111 let groups = data_count.div_ceil(data_n);
112 for g in 0..groups {
113 let data_start = g * data_n;
114 let data_end = (data_start + data_n).min(data_count);
115 let parity_start = data_count + g * parity_n;
116 let mut have = 0usize;
117 for i in data_start..data_end {
118 if present.get(i).copied() == Some(true) {
119 have += 1;
120 }
121 }
122 have += data_n - (data_end - data_start);
124 for p in 0..parity_n {
125 if present.get(parity_start + p).copied() == Some(true) {
126 have += 1;
127 }
128 }
129 if have < data_n {
130 return false;
131 }
132 let _ = stripe;
133 }
134 true
135}
136
137pub fn reconstruct(manifest: &Manifest, slots: &[Option<Vec<u8>>]) -> Result<Vec<Vec<u8>>> {
139 let Some(info) = &manifest.erasure else {
140 let mut out = Vec::new();
141 for (i, refer) in manifest.chunks.iter().enumerate() {
142 let Some(data) = slots.get(i).and_then(|s| s.as_ref()) else {
143 return Err(DdError::protocol(
144 ErrorCode::Dds2003MissingChunk,
145 format!("chunk {i}"),
146 ));
147 };
148 super::verify_chunk(&refer.id, data)?;
149 out.push(data.clone());
150 }
151 return Ok(out);
152 };
153 let data_n = info.spec.data_shards as usize;
154 let parity_n = info.spec.parity_shards as usize;
155 let data_count = info.data_count as usize;
156 let groups = data_count.div_ceil(data_n);
157 let mut recovered = vec![Vec::new(); data_count];
158 for g in 0..groups {
159 let data_start = g * data_n;
160 let data_end = (data_start + data_n).min(data_count);
161 let parity_start = data_count + g * parity_n;
162 let width = (data_start..data_end)
163 .filter_map(|i| slots.get(i).and_then(|s| s.as_ref()).map(|d| d.len()))
164 .chain((0..parity_n).filter_map(|p| {
165 slots
166 .get(parity_start + p)
167 .and_then(|s| s.as_ref())
168 .map(|d| d.len())
169 }))
170 .max()
171 .unwrap_or(0);
172 let mut shards: Vec<Option<Vec<u8>>> = vec![None; data_n + parity_n];
173 for (off, i) in (data_start..data_end).enumerate() {
174 if let Some(Some(d)) = slots.get(i) {
175 let mut padded = vec![0u8; width];
176 padded[..d.len()].copy_from_slice(d);
177 shards[off] = Some(padded);
178 }
179 }
180 #[allow(clippy::needless_range_loop)]
181 for off in (data_end - data_start)..data_n {
182 shards[off] = Some(vec![0u8; width]);
183 }
184 for p in 0..parity_n {
185 if let Some(Some(d)) = slots.get(parity_start + p) {
186 shards[data_n + p] = Some(d.clone());
187 }
188 }
189 if parity_n == 1 {
190 recover_xor(&mut shards, data_n, width)?;
191 } else {
192 let rs = ReedSolomon::new(data_n, parity_n).map_err(|e| {
193 DdError::protocol(ErrorCode::Dds2002CorruptChunk, format!("reed-solomon: {e}"))
194 })?;
195 rs.reconstruct(&mut shards).map_err(|_| {
196 DdError::protocol(ErrorCode::Dds2003MissingChunk, "erasure reconstruct")
197 })?;
198 }
199 for (off, i) in (data_start..data_end).enumerate() {
200 let shard = shards[off]
201 .as_ref()
202 .ok_or_else(|| DdError::protocol(ErrorCode::Dds2003MissingChunk, "shard"))?;
203 let want = manifest.chunks[i].length as usize;
204 recovered[i] = shard[..want].to_vec();
205 super::verify_chunk(&manifest.chunks[i].id, &recovered[i])?;
206 }
207 }
208 Ok(recovered)
209}
210
211fn recover_xor(shards: &mut [Option<Vec<u8>>], _data_n: usize, width: usize) -> Result<()> {
212 let missing: Vec<usize> = shards
213 .iter()
214 .enumerate()
215 .filter(|(_, s)| s.is_none())
216 .map(|(i, _)| i)
217 .collect();
218 if missing.is_empty() {
219 return Ok(());
220 }
221 if missing.len() > 1 {
222 return Err(DdError::protocol(
223 ErrorCode::Dds2003MissingChunk,
224 "xor erasure recovers at most one shard per group",
225 ));
226 }
227 let mut acc = vec![0u8; width];
228 for (i, shard) in shards.iter().enumerate() {
229 if i == missing[0] {
230 continue;
231 }
232 let Some(s) = shard else {
233 return Err(DdError::protocol(
234 ErrorCode::Dds2003MissingChunk,
235 "xor group incomplete",
236 ));
237 };
238 for (a, b) in acc.iter_mut().zip(s.iter()) {
239 *a ^= *b;
240 }
241 }
242 shards[missing[0]] = Some(acc);
243 Ok(())
244}
245
246#[cfg(test)]
247mod tests {
248 use super::*;
249 use crate::chunk::{chunk_payload, default_fixed, reassemble};
250
251 #[test]
252 fn xor_recovers_one_missing() {
253 let data = vec![9u8; 90_000];
254 let raw = chunk_payload(&data, default_fixed()).unwrap();
255 let body = apply_erasure(
256 raw,
257 ErasureSpec {
258 data_shards: 2,
259 parity_shards: 1,
260 },
261 )
262 .unwrap();
263 let mut slots: Vec<Option<Vec<u8>>> = body.chunks.iter().cloned().map(Some).collect();
264 slots[0] = None;
265 let got = reconstruct(&body.manifest, &slots).unwrap();
266 let out = reassemble(
267 &Manifest {
268 erasure: None,
269 chunks: body.manifest.chunks
270 [..body.manifest.erasure.as_ref().unwrap().data_count as usize]
271 .to_vec(),
272 ..body.manifest.clone()
273 },
274 &got,
275 )
276 .unwrap();
277 assert_eq!(out, data);
278 }
279
280 #[test]
281 fn rs_recovers_two_missing() {
282 let data = vec![3u8; 200_000];
283 let raw = chunk_payload(&data, default_fixed()).unwrap();
284 let body = apply_erasure(
285 raw,
286 ErasureSpec {
287 data_shards: 3,
288 parity_shards: 2,
289 },
290 )
291 .unwrap();
292 let mut slots: Vec<Option<Vec<u8>>> = body.chunks.iter().cloned().map(Some).collect();
293 slots[0] = None;
294 slots[1] = None;
295 let got = reconstruct(&body.manifest, &slots).unwrap();
296 assert_eq!(got.iter().map(|c| c.len()).sum::<usize>(), data.len());
297 let mut cat = Vec::new();
298 for c in got {
299 cat.extend_from_slice(&c);
300 }
301 assert_eq!(cat, data);
302 }
303}