Skip to main content

rtc_dtls/crypto/
padding.rs

1use cbc::cipher::block_padding::{PadType, RawPadding, UnpadError};
2use core::panic;
3
4/// DTLS block-cipher padding, as a marker type for the padding scheme.
5///
6/// Has no values — it exists to parameterize the CBC cipher over its padding.
7pub enum DtlsPadding {}
8/// Reference: RFC5246, 6.2.3.2
9impl RawPadding for DtlsPadding {
10    const TYPE: PadType = PadType::Reversible;
11
12    fn raw_pad(block: &mut [u8], pos: usize) {
13        if pos >= block.len() {
14            panic!("`pos` is bigger or equal to block size");
15        }
16
17        let padding_length = block.len() - pos - 1;
18        if padding_length > 255 {
19            panic!("block size is too big for DTLS");
20        }
21
22        set(&mut block[pos..], padding_length as u8);
23    }
24
25    fn raw_unpad(data: &[u8]) -> Result<&[u8], UnpadError> {
26        let padding_length = data.last().copied().unwrap_or(1) as usize;
27        if padding_length + 1 > data.len() {
28            return Err(UnpadError);
29        }
30
31        let padding_begin = data.len() - padding_length - 1;
32
33        if data[padding_begin..data.len() - 1]
34            .iter()
35            .any(|&byte| byte as usize != padding_length)
36        {
37            return Err(UnpadError);
38        }
39
40        Ok(&data[0..padding_begin])
41    }
42}
43
44/// Sets all bytes in `dst` equal to `value`
45#[inline(always)]
46fn set(dst: &mut [u8], value: u8) {
47    // SAFETY: we overwrite valid memory behind `dst`
48    // note: loop is not used here because it produces
49    // unnecessary branch which tests for zero-length slices
50    unsafe {
51        core::ptr::write_bytes(dst.as_mut_ptr(), value, dst.len());
52    }
53}
54
55#[cfg(test)]
56mod tests {
57    use rand::RngExt;
58
59    use super::*;
60
61    #[test]
62    fn padding_length_is_amount_of_bytes_excluding_the_padding_length_itself() -> Result<(), ()> {
63        for original_length in 0..128 {
64            for padding_length in 0..(256 - original_length) {
65                let mut block = vec![0; original_length + padding_length + 1];
66                rand::rng().fill(&mut block[0..original_length]);
67                let original = block[0..original_length].to_vec();
68                DtlsPadding::raw_pad(&mut block, original_length);
69
70                for byte in block[original_length..].iter() {
71                    assert_eq!(*byte as usize, padding_length);
72                }
73                assert_eq!(block[0..original_length], original);
74            }
75        }
76
77        Ok(())
78    }
79
80    #[test]
81    #[should_panic]
82    fn full_block_is_padding_error() {
83        for original_length in 0..256 {
84            let mut block = vec![0; original_length];
85            DtlsPadding::raw_pad(&mut block, original_length);
86        }
87    }
88
89    #[test]
90    #[should_panic]
91    fn padding_length_bigger_than_255_is_a_pad_error() {
92        let padding_length = 256;
93        for original_length in 0..128 {
94            let mut block = vec![0; original_length + padding_length + 1];
95            DtlsPadding::raw_pad(&mut block, original_length);
96        }
97    }
98
99    #[test]
100    fn empty_block_is_unpadding_error() {
101        let r = DtlsPadding::raw_unpad(&[]);
102        assert!(r.is_err());
103    }
104
105    #[test]
106    fn padding_too_big_for_block_is_unpadding_error() {
107        let r = DtlsPadding::raw_unpad(&[1]);
108        assert!(r.is_err());
109    }
110
111    #[test]
112    fn one_of_the_padding_bytes_with_value_different_than_padding_length_is_unpadding_error() {
113        for padding_length in 0..16 {
114            for invalid_byte in 0..padding_length {
115                let mut block = vec![0; padding_length + 1];
116                DtlsPadding::raw_pad(&mut block, 0);
117
118                assert_eq!(DtlsPadding::raw_unpad(&block).ok(), Some(&[][..]));
119                block[invalid_byte] = (padding_length - 1) as u8;
120                let r = DtlsPadding::raw_unpad(&block);
121                assert!(r.is_err());
122            }
123        }
124    }
125}