Skip to main content

eggress_protocol_shadowsocks/
nonce.rs

1use crate::ShadowsocksError;
2
3#[derive(Debug, Clone)]
4pub struct NonceCounter {
5    nonce_size: usize,
6    counter: u64,
7}
8
9impl NonceCounter {
10    pub fn new(nonce_size: usize) -> Self {
11        Self {
12            nonce_size,
13            counter: 0,
14        }
15    }
16
17    pub fn starting_at(nonce_size: usize, counter: u64) -> Self {
18        Self {
19            nonce_size,
20            counter,
21        }
22    }
23
24    pub fn current(&self, buf: &mut [u8]) -> Result<(), ShadowsocksError> {
25        if buf.len() != self.nonce_size {
26            return Err(ShadowsocksError::Other("nonce size mismatch".into()));
27        }
28        buf.fill(0);
29        let end = self.nonce_size.min(8);
30        buf[..end].copy_from_slice(&self.counter.to_le_bytes()[..end]);
31        Ok(())
32    }
33
34    pub fn advance(&mut self) -> Result<(), ShadowsocksError> {
35        self.counter = self
36            .counter
37            .checked_add(1)
38            .ok_or_else(|| ShadowsocksError::Other("nonce counter overflow".into()))?;
39        Ok(())
40    }
41
42    pub fn nonce_size(&self) -> usize {
43        self.nonce_size
44    }
45}
46
47#[cfg(test)]
48mod tests {
49    use super::*;
50
51    #[test]
52    fn test_nonce_starts_at_zero() {
53        let nonce = NonceCounter::new(12);
54        let mut bytes = [0u8; 12];
55        nonce.current(&mut bytes).unwrap();
56        assert_eq!(bytes, [0u8; 12]);
57    }
58
59    #[test]
60    fn test_nonce_increments() {
61        let mut nonce = NonceCounter::new(12);
62        nonce.advance().unwrap();
63        let mut bytes = [0u8; 12];
64        nonce.current(&mut bytes).unwrap();
65        assert_eq!(bytes.len(), 12);
66        // little-endian: counter in first 8 bytes, rest zero
67        assert_eq!(&bytes[..8], &1u64.to_le_bytes());
68        assert_eq!(&bytes[8..], &[0, 0, 0, 0]);
69    }
70
71    #[test]
72    fn test_nonce_advance_multiple() {
73        let mut nonce = NonceCounter::new(12);
74        for i in 1u64..=10 {
75            nonce.advance().unwrap();
76            let mut bytes = [0u8; 12];
77            nonce.current(&mut bytes).unwrap();
78            assert_eq!(&bytes[..8], &i.to_le_bytes());
79        }
80    }
81
82    #[test]
83    fn test_nonce_overflow_returns_error() {
84        let mut nonce = NonceCounter::new(12);
85        nonce.counter = u64::MAX;
86        assert!(nonce.advance().is_err());
87    }
88}