Skip to main content

dht_rpc/
cenc.rs

1use std::{
2    convert::TryFrom,
3    net::{Ipv4Addr, SocketAddr},
4};
5
6use compact_encoding::{
7    CompactEncoding, EncodingError, VecEncodable, decode_usize, encode_usize_var,
8    encoded_size_usize, map_decode, take_array, vec_encoded_size_for_fixed_sized_elements,
9    write_array,
10};
11
12use crate::{
13    Command, Error, ExternalCommand, IdBytes, InternalCommand, Peer, Result,
14    constants::{HASH_SIZE, ID_SIZE, REQUEST_ID, RESPONSE_ID},
15    message::{MsgData, ReplyMsgData, RequestMsgData},
16    socket_into_v4,
17};
18
19impl CompactEncoding for InternalCommand {
20    fn encoded_size(&self) -> std::result::Result<usize, EncodingError> {
21        Ok(1)
22    }
23
24    fn encode<'a>(&self, buffer: &'a mut [u8]) -> std::result::Result<&'a mut [u8], EncodingError> {
25        write_array(&[*self as u8], buffer)
26    }
27
28    fn decode(buffer: &[u8]) -> std::result::Result<(Self, &[u8]), EncodingError>
29    where
30        Self: Sized,
31    {
32        let ([value], rest) = take_array::<1>(buffer)?;
33        let cmd = InternalCommand::try_from(value).map_err(EncodingError::from)?;
34        Ok((cmd, rest))
35    }
36}
37
38impl From<Error> for EncodingError {
39    fn from(value: Error) -> Self {
40        EncodingError {
41            kind: compact_encoding::EncodingErrorKind::InvalidData,
42            message: value.to_string(),
43        }
44    }
45}
46
47impl CompactEncoding for Peer {
48    fn encoded_size(&self) -> std::result::Result<usize, EncodingError> {
49        Ok(6)
50    }
51
52    fn encode<'a>(&self, buffer: &'a mut [u8]) -> std::result::Result<&'a mut [u8], EncodingError> {
53        let ip = self.socketv4()?;
54        ip.encode(buffer)
55    }
56
57    fn decode(buffer: &[u8]) -> std::result::Result<(Self, &[u8]), EncodingError>
58    where
59        Self: Sized,
60    {
61        let ((ip, port), rest) = map_decode!(buffer, [Ipv4Addr, u16]);
62        Ok((
63            Peer {
64                id: None,
65                addr: SocketAddr::from((ip, port)),
66                referrer: None,
67            },
68            rest,
69        ))
70    }
71}
72
73impl VecEncodable for Peer {
74    fn vec_encoded_size(vec: &[Self]) -> std::result::Result<usize, EncodingError>
75    where
76        Self: Sized,
77    {
78        Ok(vec_encoded_size_for_fixed_sized_elements(
79            vec,
80            Peer::ENCODED_SIZE,
81        ))
82    }
83}
84
85const IP_AND_PORT_NUM_BYTES: usize = 6;
86
87/// TODO this will panic for ipv6
88fn id_from_socket(addr: &SocketAddr) -> [u8; ID_SIZE] {
89    let addr = socket_into_v4(addr).expect("TODO panics for ipv6");
90    let mut from_buff = vec![0; IP_AND_PORT_NUM_BYTES];
91    addr.encode(&mut from_buff).expect("should always fit");
92
93    generic_hash(&from_buff)
94}
95
96pub(crate) fn calculate_peer_id(from: &Peer) -> [u8; ID_SIZE] {
97    id_from_socket(&from.addr)
98}
99
100#[allow(unsafe_code, reason = "needed for libsodium bindings")]
101pub fn generic_hash(input: &[u8]) -> [u8; HASH_SIZE] {
102    let mut out = [0; HASH_SIZE];
103    let ret = unsafe {
104        libsodium_sys::crypto_generichash(
105            out.as_mut_ptr(),
106            out.len(),
107            input.as_ptr(),
108            input.len() as u64,
109            std::ptr::null(),
110            0,
111        )
112    };
113    if ret != 0 {
114        panic!("Only errors when the input is invalid. Inputs here or checked");
115    }
116    out
117}
118
119#[allow(unsafe_code, reason = "needed for libsodium bindings")]
120pub(crate) fn generic_hash_with_key(input: &[u8], key: &[u8]) -> Result<[u8; HASH_SIZE]> {
121    let mut out = [0; HASH_SIZE];
122    let ret = unsafe {
123        libsodium_sys::crypto_generichash(
124            out.as_mut_ptr(),
125            out.len(),
126            input.as_ptr(),
127            input.len() as u64,
128            key.as_ptr(),
129            key.len(),
130        )
131    };
132    if ret != 0 {
133        return Err(Error::LibSodiumGenericHashError(ret));
134    }
135    Ok(out)
136}
137
138pub(crate) fn validate_id(id: &Option<[u8; ID_SIZE]>, from: &Peer) -> Option<IdBytes> {
139    if let Some(id) = id
140        && id == &calculate_peer_id(from)
141    {
142        return Some(IdBytes::from(*id));
143    }
144    None
145}
146
147macro_rules! maybe_add_flag {
148    ($cond:expr, $shift:expr) => {
149        if $cond { 1 << $shift } else { 0 }
150    };
151}
152
153macro_rules! maybe_decode {
154    ($type:ty, $cond:expr, $buf:expr) => {
155        if $cond {
156            let (out, rest) = <$type>::decode($buf)?;
157            (Some(out), rest)
158        } else {
159            (None, $buf)
160        }
161    };
162}
163
164impl CompactEncoding for RequestMsgData {
165    fn encoded_size(&self) -> std::result::Result<usize, EncodingError> {
166        let mut out = 1 + // REQUEST_ID
167                      1 + // flags
168                      2 + // tid
169                      6 + // peer
170                      1   // command byte
171                    ;
172        if self.id.is_some() {
173            out += ID_SIZE;
174        }
175        if self.token.is_some() {
176            out += 32;
177        }
178        if self.target.is_some() {
179            out += 32;
180        }
181        if let Some(v) = &self.value {
182            out += v.encoded_size()?;
183        }
184        Ok(out)
185    }
186
187    fn encode<'a>(&self, buffer: &'a mut [u8]) -> std::result::Result<&'a mut [u8], EncodingError> {
188        let mut flags: u8 = 0;
189        let is_internal = matches!(self.command, Command::Internal(_));
190        flags |= maybe_add_flag!(self.id.is_some(), 0);
191        flags |= maybe_add_flag!(self.token.is_some(), 1);
192        flags |= maybe_add_flag!(is_internal, 2);
193        flags |= maybe_add_flag!(self.target.is_some(), 3);
194        flags |= maybe_add_flag!(self.value.is_some(), 4);
195
196        let mut rest = write_array(&[REQUEST_ID, flags], buffer)?;
197        rest = self.tid.encode(rest)?;
198        rest = CompactEncoding::encode(&self.to, rest)?;
199        if let Some(id) = &self.id {
200            rest = id.encode(rest)?;
201        }
202        if let Some(token) = &self.token {
203            rest = token.encode(rest)?;
204        }
205        rest = u8::encode(&self.command.encode(), rest)?;
206        if let Some(target) = &self.target {
207            rest = target.encode(rest)?;
208        }
209        if let Some(v) = &self.value {
210            rest = v.encode(rest)?
211        }
212        //println!(
213        //    "
214        //MSGENCODE
215        //etid = {}
216        //einternal = {}
217        //ecommand = {}
218        //",
219        //    self.tid, is_internal, self.command
220        //);
221        Ok(rest)
222    }
223
224    fn decode(buffer: &[u8]) -> std::result::Result<(Self, &[u8]), EncodingError>
225    where
226        Self: Sized,
227    {
228        let (([_req_flag, flags], tid, to), rest) = map_decode!(buffer, [[u8; 2], u16, Peer]);
229        // assert_eq!(_req_flag, REQUEST_ID)
230        let (id, rest) = maybe_decode!([u8; 32], flags & (1 << 0) != 0, rest);
231        let (token, rest) = maybe_decode!([u8; 32], flags & (1 << 1) != 0, rest);
232        let internal = (flags & 1 << 2) != 0;
233        let ([cmd_u8], rest) = take_array::<1>(rest)?;
234        let command = if internal {
235            Command::from(InternalCommand::try_from(cmd_u8).map_err(EncodingError::from)?)
236        } else {
237            Command::from(ExternalCommand(cmd_u8 as usize))
238        };
239        let (target, rest) = maybe_decode!([u8; 32], flags & (1 << 3) != 0, rest);
240        let (value, rest) = maybe_decode!(Vec<u8>, flags & (1 << 4) != 0, rest);
241        Ok((
242            Self {
243                tid,
244                to,
245                id,
246                token,
247                command,
248                target,
249                value,
250            },
251            rest,
252        ))
253    }
254}
255
256impl CompactEncoding for ReplyMsgData {
257    fn encoded_size(&self) -> std::result::Result<usize, EncodingError> {
258        let mut out: usize = 1 + // RESPONSE_ID
259                             1 + // flags
260                             6 + // to
261                             2   // tid
262                            ;
263        if self.id.is_some() {
264            out += 32;
265        }
266        if self.token.is_some() {
267            out += 32;
268        }
269        if !self.closer_nodes.is_empty() {
270            out += self.closer_nodes.encoded_size()?;
271        }
272        if self.error > 0 {
273            out += encoded_size_usize(self.error);
274        }
275        if let Some(v) = &self.value {
276            out += v.encoded_size()?;
277        }
278
279        Ok(out)
280    }
281
282    fn encode<'a>(&self, buffer: &'a mut [u8]) -> std::result::Result<&'a mut [u8], EncodingError> {
283        let mut flags: u8 = 0;
284        flags |= maybe_add_flag!(self.id.is_some(), 0);
285        flags |= maybe_add_flag!(self.token.is_some(), 1);
286        flags |= maybe_add_flag!(!self.closer_nodes.is_empty(), 2);
287        flags |= maybe_add_flag!(self.error > 0, 3);
288        flags |= maybe_add_flag!(self.value.is_some(), 4);
289
290        let mut rest = write_array(&[RESPONSE_ID, flags], buffer)?;
291        rest = self.tid.encode(rest)?;
292        rest = CompactEncoding::encode(&self.to, rest)?;
293        if let Some(id) = &self.id {
294            rest = id.encode(rest)?;
295        }
296        if let Some(token) = &self.token {
297            rest = token.encode(rest)?;
298        }
299        if !self.closer_nodes.is_empty() {
300            rest = self.closer_nodes.encode(rest)?;
301        }
302        if self.error > 0 {
303            rest = encode_usize_var(&self.error, rest)?;
304        }
305        if let Some(v) = &self.value {
306            rest = v.encode(rest)?
307        }
308        Ok(rest)
309    }
310
311    fn decode(buffer: &[u8]) -> std::result::Result<(Self, &[u8]), EncodingError>
312    where
313        Self: Sized,
314    {
315        let (([_, flags], tid, to), rest) = map_decode!(buffer, [[u8; 2], u16, Peer]);
316        let (id, rest) = maybe_decode!([u8; 32], flags & (1 << 0) != 0, rest);
317        let (token, rest) = maybe_decode!([u8; 32], flags & (1 << 1) != 0, rest);
318        let (closer_nodes, rest) = if flags & (1 << 2) != 0 {
319            <Vec<Peer> as CompactEncoding>::decode(rest)?
320        } else {
321            (vec![], rest)
322        };
323        let (error, rest) = if flags & (1 << 3) != 0 {
324            decode_usize(rest)?
325        } else {
326            (0, rest)
327        };
328        let (value, rest) = maybe_decode!(Vec<u8>, flags & (1 << 4) != 0, rest);
329        Ok((
330            Self {
331                tid,
332                to,
333                id,
334                token,
335                closer_nodes,
336                error,
337                value,
338            },
339            rest,
340        ))
341    }
342}
343
344impl CompactEncoding for MsgData {
345    fn encoded_size(&self) -> std::result::Result<usize, EncodingError> {
346        match self {
347            MsgData::Request(x) => x.encoded_size(),
348            MsgData::Reply(x) => x.encoded_size(),
349        }
350    }
351
352    fn encode<'a>(&self, buffer: &'a mut [u8]) -> std::result::Result<&'a mut [u8], EncodingError> {
353        match self {
354            MsgData::Request(x) => x.encode(buffer),
355            MsgData::Reply(x) => x.encode(buffer),
356        }
357    }
358
359    fn decode(buffer: &[u8]) -> std::result::Result<(Self, &[u8]), EncodingError>
360    where
361        Self: Sized,
362    {
363        let req_resp_flag = buffer[0];
364        Ok(match req_resp_flag {
365            REQUEST_ID => {
366                let (msg, rest) = RequestMsgData::decode(buffer)?;
367                (MsgData::Request(msg), rest)
368            }
369            RESPONSE_ID => {
370                let (msg, rest) = ReplyMsgData::decode(buffer)?;
371                (MsgData::Reply(msg), rest)
372            }
373            _ => {
374                return Err(EncodingError::invalid_data(&format!(
375                    "Could not decode MsgData. The first byte [{req_resp_flag}] did not match the request [{REQUEST_ID}] or response [{RESPONSE_ID}] flags"
376                )));
377            }
378        })
379    }
380}