Skip to main content

deaddrop_core/chunk/
erasure.rs

1//! Optional forward error correction for Drop chunks.
2//!
3//! `erasure-xor-v1` (parity_shards == 1) uses per-group XOR.
4//! `erasure-v1` (parity_shards > 1) uses Reed–Solomon over GF(2^8).
5//! Recovered payload is still verified against `payload_hash`.
6
7use 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    /// Number of original payload chunks (prefix of `manifest.chunks`).
23    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
44/// Append parity chunks. Data chunks stay content-addressed and independently verifiable.
45pub 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        // Empty padding shards (incomplete last group) count as present.
123        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
137/// Reconstruct original data chunks. `slots[i]` is `None` when missing.
138pub 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}