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
87fn 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 + 1 + 2 + 6 + 1 ;
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 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 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 + 1 + 6 + 2 ;
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}