massa_bootstrap/bindings/
server.rs1use crate::bindings::BindingReadExact;
4use crate::error::BootstrapError;
5use crate::messages::{
6 BootstrapClientMessage, BootstrapClientMessageDeserializer, BootstrapServerMessage,
7 BootstrapServerMessageSerializer,
8};
9use crate::settings::BootstrapSrvBindCfg;
10use massa_hash::Hash;
11use massa_hash::HASH_SIZE_BYTES;
12use massa_models::config::{
13 BOOTSTRAP_MESSAGE_LEN_PREFIX_MAX, BOOTSTRAP_MESSAGE_LEN_PREFIX_SIZE_BYTES,
14 MAX_BOOTSTRAP_MESSAGE_FROM_CLIENT_SIZE, MAX_BOOTSTRAP_MESSAGE_FROM_SERVER_SIZE,
15};
16use massa_models::serialization::{DeserializeMinBEInt, SerializeMinBEInt};
17use massa_models::version::{Version, VersionDeserializer, VersionSerializer};
18use massa_serialization::{DeserializeError, Deserializer, Serializer};
19use massa_signature::KeyPair;
20use massa_time::MassaTime;
21use std::io;
22use std::sync::atomic::{AtomicUsize, Ordering};
23use std::time::Instant;
24use std::{
25 convert::TryInto,
26 io::ErrorKind,
27 net::{SocketAddr, TcpStream},
28 thread,
29 time::Duration,
30};
31use stream_limiter::{Limiter, LimiterOptions};
32use tracing::{error, warn};
33
34use super::BindingWriteExact;
35
36const MAX_CONCURRENT_ERROR_SENDS: usize = 32;
44
45static ONGOING_ERROR_SENDS: AtomicUsize = AtomicUsize::new(0);
48
49struct ErrorSendSlot;
52
53impl ErrorSendSlot {
54 fn claim() -> Option<Self> {
56 ONGOING_ERROR_SENDS
57 .fetch_update(Ordering::AcqRel, Ordering::Acquire, |ongoing| {
58 (ongoing < MAX_CONCURRENT_ERROR_SENDS).then_some(ongoing + 1)
59 })
60 .ok()
61 .map(|_| ErrorSendSlot)
62 }
63}
64
65impl Drop for ErrorSendSlot {
66 fn drop(&mut self) {
67 ONGOING_ERROR_SENDS.fetch_sub(1, Ordering::AcqRel);
68 }
69}
70
71const KNOWN_PREFIX_FROM_CLIENT_LEN: usize =
72 HASH_SIZE_BYTES + BOOTSTRAP_MESSAGE_LEN_PREFIX_SIZE_BYTES;
73struct ClientMessageLeader {
75 received_prev_hash: Option<Hash>,
76 msg_len: u32,
77}
78
79pub struct BootstrapServerBinder {
81 max_consensus_block_ids: u64,
83 thread_count: u8,
84 max_datastore_key_length: u8,
85 randomness_size_bytes: usize,
86 local_keypair: KeyPair,
87 duplex: Limiter<TcpStream>,
88 prev_message: Option<Hash>,
89 version_serializer: VersionSerializer,
90 version_deserializer: VersionDeserializer,
91 write_error_timeout: MassaTime,
92}
93
94impl BootstrapServerBinder {
95 #[allow(clippy::too_many_arguments)]
102 pub fn new(
103 duplex: TcpStream,
104 local_keypair: KeyPair,
105 cfg: BootstrapSrvBindCfg,
106 rw_limit: Option<u64>,
107 ) -> Self {
108 let BootstrapSrvBindCfg {
109 rate_limit: _limit,
110 thread_count,
111 max_datastore_key_length,
112 randomness_size_bytes,
113 consensus_bootstrap_part_size: _part_size,
114 max_consensus_block_ids,
115 write_error_timeout,
116 } = cfg;
117
118 let limit_opts = rw_limit.map(|limit| -> LimiterOptions {
119 LimiterOptions::new(limit, Duration::from_millis(1000), limit)
120 });
121 let duplex = Limiter::new(duplex, limit_opts.clone(), limit_opts);
122 BootstrapServerBinder {
123 max_consensus_block_ids,
124 local_keypair,
125 duplex,
126 prev_message: None,
127 thread_count,
128 max_datastore_key_length,
129 randomness_size_bytes,
130 version_serializer: VersionSerializer::new(),
131 version_deserializer: VersionDeserializer::new(),
132 write_error_timeout,
133 }
134 }
135 pub fn handshake_timeout(
138 &mut self,
139 version: Version,
140 duration: Option<Duration>,
141 ) -> Result<(), BootstrapError> {
142 let deadline = duration.map(|d| Instant::now() + d);
143 let msg_hash = {
145 let mut version_bytes = Vec::new();
146 self.version_serializer
147 .serialize(&version, &mut version_bytes)?;
148 let mut msg_bytes = vec![0u8; version_bytes.len() + self.randomness_size_bytes];
149 self.read_exact_timeout(&mut msg_bytes, deadline)
150 .map_err(|(e, _)| e)?;
151 let (_, received_version) = self
152 .version_deserializer
153 .deserialize::<DeserializeError>(&msg_bytes[..version_bytes.len()])
154 .map_err(|err| BootstrapError::GeneralError(format!("{}", &err)))?;
155 if !version.is_compatible(&received_version) {
156 return Err(BootstrapError::IncompatibleVersionError(format!("Received a bad incompatible version in handshake. (excepted: {}, received: {})", version, received_version)));
157 }
158 Hash::compute_from(&msg_bytes)
159 };
160
161 self.prev_message = Some(msg_hash);
163
164 Ok(())
165 }
166
167 pub fn send_msg(
168 &mut self,
169 timeout: Duration,
170 msg: BootstrapServerMessage,
171 ) -> Result<(), BootstrapError> {
172 let to_str = msg.to_string();
173 self.send_timeout(msg, Some(timeout)).map_err(|e| match e {
174 BootstrapError::IoError(e)
175 if e.kind() == ErrorKind::TimedOut || e.kind() == ErrorKind::WouldBlock =>
177 {
178 BootstrapError::TimedOut(std::io::Error::new(
179 std::io::ErrorKind::TimedOut,
180 format!("BootstrapServerMessage::{} send timed out", to_str),
181 ))
182 }
183 _ => e,
184 })
185 }
186
187 pub(crate) fn close_and_send_error<F>(mut self, msg: String, addr: SocketAddr, close_fn: F)
197 where
198 F: FnOnce() + Send + 'static,
199 {
200 let Some(slot) = ErrorSendSlot::claim() else {
203 warn!(
204 "bootstrap server closing connection from {} without sending error '{}': too many refusals in flight",
205 addr, msg
206 );
207 return;
208 };
209
210 let spawned = thread::Builder::new()
211 .name("bootstrap-error-send".to_string())
212 .spawn(move || {
213 let _slot = slot;
215 let msg_cloned = msg.clone();
216 let err_send = self.send_error_timeout(msg_cloned);
217 match err_send {
218 Err(BootstrapError::IoError(e)) if e.kind() == ErrorKind::TimedOut => error!(
219 "bootstrap server timed out sending error '{}' to addr {}",
220 msg, addr
221 ),
222 Err(e) => error!(
223 "bootstrap server encountered error '{}' sending error '{}' to addr '{}'",
224 e, msg, addr
225 ),
226 Ok(_) => {}
227 }
228 close_fn();
229 });
230
231 if let Err(e) = spawned {
235 error!(
236 "bootstrap server failed to spawn the error-send thread for addr {}: {}",
237 addr, e
238 );
239 }
240 }
241 pub fn send_error_timeout(&mut self, error: String) -> Result<(), BootstrapError> {
242 self.send_timeout(
243 BootstrapServerMessage::BootstrapError { error },
244 Some(self.write_error_timeout.to_duration()),
245 )
246 .map_err(|e| match e {
247 BootstrapError::IoError(e) if e.kind() == ErrorKind::WouldBlock => {
248 BootstrapError::TimedOut(ErrorKind::TimedOut.into())
249 }
250 e => e,
251 })
252 }
253
254 pub fn send_timeout(
256 &mut self,
257 msg: BootstrapServerMessage,
258 duration: Option<Duration>,
259 ) -> Result<(), BootstrapError> {
260 let deadline = duration.map(|d| Instant::now() + d);
261 let mut msg_bytes = Vec::new();
263 BootstrapServerMessageSerializer::new().serialize(&msg, &mut msg_bytes)?;
264 let msg_len: u32 = msg_bytes.len().try_into().map_err(|e| {
265 BootstrapError::GeneralError(format!("bootstrap message too large to encode: {}", e))
266 })?;
267
268 let sig = {
270 if let Some(prev_message) = self.prev_message {
271 let mut signed_data =
273 Vec::with_capacity(HASH_SIZE_BYTES.saturating_add(msg_len as usize));
274 signed_data.extend(prev_message.to_bytes());
275 signed_data.extend(&msg_bytes);
276 self.local_keypair.sign(&Hash::compute_from(&signed_data))?
277 } else {
278 self.local_keypair.sign(&Hash::compute_from(&msg_bytes))?
280 }
281 };
282
283 let msg_len_bytes = msg_len.to_be_bytes_min(MAX_BOOTSTRAP_MESSAGE_FROM_SERVER_SIZE)?;
285
286 let stream_data = [sig.to_bytes().as_slice(), &msg_len_bytes, &msg_bytes].concat();
288
289 self.write_all_timeout(&stream_data, deadline)
291 .map_err(|(e, _)| e)?;
292
293 self.prev_message = Some(Hash::compute_from(&sig.to_bytes()));
295
296 Ok(())
297 }
298
299 pub fn next_timeout(
302 &mut self,
303 duration: Option<Duration>,
304 ) -> Result<BootstrapClientMessage, BootstrapError> {
305 let deadline = duration.map(|d| Instant::now() + d);
306
307 let mut known_len_buf = vec![0; KNOWN_PREFIX_FROM_CLIENT_LEN];
308 self.read_exact_timeout(&mut known_len_buf, deadline)
310 .map_err(|(err, _consumed)| err)?;
311
312 let ClientMessageLeader {
313 received_prev_hash,
314 msg_len,
315 } = self.decode_message_leader(&known_len_buf)?;
316
317 let mut msg_bytes = vec![0u8; msg_len as usize];
319 self.read_exact_timeout(&mut msg_bytes, deadline)
320 .map_err(|(err, _consumed)| err)?;
321
322 if received_prev_hash != self.prev_message {
324 return Err(BootstrapError::GeneralError(
325 "Message sequencing has been broken".to_string(),
326 ));
327 }
328
329 if let Some(prev_hash) = received_prev_hash {
331 let mut hashed_bytes =
333 Vec::with_capacity(HASH_SIZE_BYTES.saturating_add(msg_bytes.len()));
334 hashed_bytes.extend(prev_hash.to_bytes());
335 hashed_bytes.extend(&msg_bytes);
336 self.prev_message = Some(Hash::compute_from(&hashed_bytes));
337 } else {
338 self.prev_message = Some(Hash::compute_from(&msg_bytes));
340 }
341
342 let (rest, msg) = BootstrapClientMessageDeserializer::new(
344 self.thread_count,
345 self.max_datastore_key_length,
346 self.max_consensus_block_ids,
347 )
348 .deserialize::<DeserializeError>(&msg_bytes)
349 .map_err(|err| BootstrapError::GeneralError(format!("{}", err)))?;
350 if !rest.is_empty() {
351 return Err(BootstrapError::GeneralError(
352 "bootstrap client message has trailing bytes after deserialization".into(),
353 ));
354 }
355
356 Ok(msg)
357 }
358
359 fn decode_message_leader(
362 &self,
363 leader_buf: &[u8],
364 ) -> Result<ClientMessageLeader, BootstrapError> {
365 let received_prev_hash = {
367 if self.prev_message.is_some() {
368 Some(Hash::from_bytes(
369 leader_buf[..HASH_SIZE_BYTES]
370 .try_into()
371 .expect("bad slice logic"),
372 ))
373 } else {
374 None
375 }
376 };
377
378 let msg_len = u32::from_be_bytes_min(
382 &leader_buf[HASH_SIZE_BYTES..],
383 BOOTSTRAP_MESSAGE_LEN_PREFIX_MAX,
384 )?
385 .0;
386 if msg_len > MAX_BOOTSTRAP_MESSAGE_FROM_CLIENT_SIZE {
387 return Err(BootstrapError::GeneralError(format!(
388 "client announced a message of {} bytes, over the {} bytes limit",
389 msg_len, MAX_BOOTSTRAP_MESSAGE_FROM_CLIENT_SIZE
390 )));
391 }
392 Ok(ClientMessageLeader {
393 received_prev_hash,
394 msg_len,
395 })
396 }
397}
398
399#[cfg(test)]
400impl BootstrapServerBinder {
401 pub(crate) fn test_prev_message_hash(&self) -> Option<Hash> {
403 self.prev_message
404 }
405}
406
407impl io::Read for BootstrapServerBinder {
408 fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
409 self.duplex.read(buf)
410 }
411}
412
413impl crate::bindings::BindingReadExact for BootstrapServerBinder {
414 fn set_read_timeout(&mut self, duration: Option<Duration>) -> Result<(), std::io::Error> {
415 if let Some(ref mut opts) = self.duplex.read_opt {
416 opts.timeout = duration;
417 }
418 self.duplex.stream.set_read_timeout(duration)
419 }
420}
421
422impl io::Write for BootstrapServerBinder {
423 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
424 self.duplex.write(buf)
425 }
426
427 fn flush(&mut self) -> io::Result<()> {
428 self.duplex.flush()
429 }
430}
431
432impl crate::bindings::BindingWriteExact for BootstrapServerBinder {
433 fn set_write_timeout(&mut self, duration: Option<Duration>) -> Result<(), std::io::Error> {
434 if let Some(ref mut opts) = self.duplex.write_opt {
435 opts.timeout = duration;
436 }
437 self.duplex.stream.set_write_timeout(duration)
438 }
439}