massa_bootstrap/bindings/
server.rs

1// Copyright (c) 2022 MASSA LABS <info@massa.net>
2
3use 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
36/// Maximum number of refusal threads running concurrently.
37///
38/// Every refused connection is told why before being closed, and that send has
39/// to happen off the bootstrap main loop so a slow client cannot stall it. A
40/// connection flood produces refusals as fast as the listener can accept, so the
41/// helper threads are capped: past the cap the socket is simply closed, which is
42/// what the client would observe on a timed-out error send anyway.
43const MAX_CONCURRENT_ERROR_SENDS: usize = 32;
44
45/// Number of refusal threads currently alive, compared against
46/// [`MAX_CONCURRENT_ERROR_SENDS`].
47static ONGOING_ERROR_SENDS: AtomicUsize = AtomicUsize::new(0);
48
49/// Releases a claimed [`ONGOING_ERROR_SENDS`] slot, panics included, so that a
50/// failing refusal can never leak the budget it took.
51struct ErrorSendSlot;
52
53impl ErrorSendSlot {
54    /// Claims a slot, or returns `None` if the budget is exhausted.
55    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;
73/// The known-length component of a message to be received.
74struct ClientMessageLeader {
75    received_prev_hash: Option<Hash>,
76    msg_len: u32,
77}
78
79/// Bootstrap server binder
80pub struct BootstrapServerBinder {
81    /// max number of block ids accepted in the client's cumulative bootstrap cursor
82    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    /// Creates a new `WriteBinder`.
96    ///
97    /// # Argument
98    /// * `duplex`: duplex stream.
99    /// * `local_keypair`: local node user keypair
100    /// * `limit`: limit max bytes per second (up and down)
101    #[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    /// Performs a handshake. Should be called after connection
136    /// MUST always be followed by a send of the `BootstrapMessage::BootstrapTime`
137    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        // read version and random bytes, send signature
144        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        // save prev sig
162        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            // On some systems, a timed out send returns WouldBlock
176                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    /// 1. Spawns a thread, unless [`MAX_CONCURRENT_ERROR_SENDS`] are already running
188    /// 2. blocks on the passed in runtime
189    /// 3. uses passed in handle to send a message to the client
190    /// 4. logs an error if the send times out
191    /// 5. runs the passed in closure (typically a custom logging msg)
192    ///
193    /// consumes the binding in the process. When the refusal budget is exhausted
194    /// or the OS refuses the spawn, the binding is dropped instead, closing the
195    /// connection without the courtesy message.
196    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        // Claim a refusal slot, or drop the connection outright rather than let
201        // a flood of refusals spawn an unbounded number of threads.
202        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                // Held for the lifetime of the thread, released on the way out.
214                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        // A failed spawn is an OS-level resource error: the slot is released with
232        // the dropped closure, and the connection closes on its own rather than
233        // taking the whole server down.
234        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    /// Writes the next message.
255    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        // serialize the message to bytes
262        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        // compute signature, and extract the bytes
269        let sig = {
270            if let Some(prev_message) = self.prev_message {
271                // there was a previous message: sign(prev_msg_hash + msg)
272                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                // there was no previous message: sign(msg)
279                self.local_keypair.sign(&Hash::compute_from(&msg_bytes))?
280            }
281        };
282
283        // construct msg length, and convert to bytes
284        let msg_len_bytes = msg_len.to_be_bytes_min(MAX_BOOTSTRAP_MESSAGE_FROM_SERVER_SIZE)?;
285
286        // organize the bytes into a sendable array
287        let stream_data = [sig.to_bytes().as_slice(), &msg_len_bytes, &msg_bytes].concat();
288
289        // send the data
290        self.write_all_timeout(&stream_data, deadline)
291            .map_err(|(e, _)| e)?;
292
293        // update prev sig
294        self.prev_message = Some(Hash::compute_from(&sig.to_bytes()));
295
296        Ok(())
297    }
298
299    // TODO: use a proper (de)serializer: https://github.com/massalabs/massa/pull/3745#discussion_r1169733161
300    /// Read a message sent from the client (not signed).
301    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        // TODO: handle a partial read
309        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        // read the rest of the message
318        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        // check previous hash
323        if received_prev_hash != self.prev_message {
324            return Err(BootstrapError::GeneralError(
325                "Message sequencing has been broken".to_string(),
326            ));
327        }
328
329        // update previous hash
330        if let Some(prev_hash) = received_prev_hash {
331            // there was a previous message: hash(prev_hash + message)
332            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            // no previous message: hash message only
339            self.prev_message = Some(Hash::compute_from(&msg_bytes));
340        }
341
342        // deserialize message
343        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    /// We are using this instead of of our library deserializer as the process is relatively straight forward
360    /// and makes error-type management cleaner
361    fn decode_message_leader(
362        &self,
363        leader_buf: &[u8],
364    ) -> Result<ClientMessageLeader, BootstrapError> {
365        // construct prev-hash
366        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        // The prefix is decoded against the wire-format bound, then the tighter client-side cap
379        // is applied to the value: the two are deliberately distinct, as the encoding width is
380        // shared with every peer while the cap is ours to tighten.
381        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    /// Post-handshake hash-chain value for tests that craft raw client frames.
402    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}