massa_bootstrap/
listener.rs1use mio::net::TcpListener;
2use mio::{Events, Interest, Poll, Token, Waker};
3use std::io::ErrorKind;
4use std::net::{SocketAddr, TcpStream};
5use std::time::Duration;
6use tracing::{info, warn};
7
8use crate::error::BootstrapError;
9use crate::tools::mio_stream_to_std;
10
11const NEW_CONNECTION: Token = Token(0);
12const STOP_LISTENER: Token = Token(10);
13
14const MAX_ACCEPTS_PER_POLL: usize = 32;
23
24pub struct BootstrapTcpListener {
26 poll: Poll,
27 events: Events,
28 server: TcpListener,
29 backlog_pending: bool,
34}
35
36pub struct BootstrapListenerStopHandle(pub(crate) Waker);
37
38fn drain_accept<S>(
49 mut accept: impl FnMut() -> std::io::Result<S>,
50 out: &mut Vec<S>,
51 max: usize,
52) -> bool {
53 for _ in 0..max {
54 match accept() {
55 Ok(item) => out.push(item),
56 Err(ref e) if e.kind() == ErrorKind::WouldBlock => return false,
57 Err(e) => {
58 warn!("Error accepting connection in bootstrap: {:?}", e);
59 return false;
60 }
61 }
62 }
63 true
64}
65
66pub enum PollEvent {
67 NewConnections(Vec<(TcpStream, SocketAddr)>),
68 Stop,
69}
70
71#[cfg_attr(test, mockall::automock)]
72impl BootstrapTcpListener {
73 pub fn create(
77 addr: &SocketAddr,
78 ) -> Result<(BootstrapListenerStopHandle, Self), BootstrapError> {
79 let domain = if addr.is_ipv4() {
80 socket2::Domain::IPV4
81 } else {
82 socket2::Domain::IPV6
83 };
84
85 let socket = socket2::Socket::new(domain, socket2::Type::STREAM, None)?;
86
87 if addr.is_ipv6() {
88 socket.set_only_v6(false)?;
89 }
90 socket.set_nonblocking(true)?;
93 socket.bind(&(*addr).into())?;
94
95 socket.listen(1024)?;
97
98 info!("Starting bootstrap listener on {}", &addr);
99 let mut server = TcpListener::from_std(socket.into());
100
101 let poll = Poll::new()?;
102
103 let waker = BootstrapListenerStopHandle(Waker::new(poll.registry(), STOP_LISTENER)?);
105
106 poll.registry()
107 .register(&mut server, NEW_CONNECTION, Interest::READABLE)?;
108
109 let events = Events::with_capacity(128);
111 Ok((
112 waker,
113 BootstrapTcpListener {
114 poll,
115 server,
116 events,
117 backlog_pending: false,
118 },
119 ))
120 }
121
122 pub fn poll(&mut self) -> Result<PollEvent, BootstrapError> {
128 let timeout = if self.backlog_pending {
132 Some(Duration::ZERO)
133 } else {
134 None
135 };
136 self.poll.poll(&mut self.events, timeout).unwrap();
137
138 let mut accept_ready = self.backlog_pending;
139 for event in self.events.iter() {
140 match event.token() {
141 NEW_CONNECTION => accept_ready = true,
142 STOP_LISTENER => {
143 return Ok(PollEvent::Stop);
144 }
145 _ => unreachable!(),
146 }
147 }
148
149 if !accept_ready {
150 return Ok(PollEvent::NewConnections(Vec::new()));
151 }
152
153 let mut accepted = Vec::with_capacity(MAX_ACCEPTS_PER_POLL);
158 self.backlog_pending =
159 drain_accept(|| self.server.accept(), &mut accepted, MAX_ACCEPTS_PER_POLL);
160
161 let mut results = Vec::with_capacity(accepted.len());
162 for (mut stream, remote_addr) in accepted {
163 let _ = self.poll.registry().deregister(&mut stream);
164 let stream: std::net::TcpStream = mio_stream_to_std(stream);
165 stream.set_nonblocking(false)?;
166 results.push((stream, remote_addr));
167 }
168
169 Ok(PollEvent::NewConnections(results))
170 }
171}
172
173impl BootstrapListenerStopHandle {
174 pub fn stop(&self) -> Result<(), BootstrapError> {
176 self.0.wake().map_err(BootstrapError::from)
177 }
178}
179
180#[cfg(test)]
181mod tests {
182 use super::{drain_accept, MAX_ACCEPTS_PER_POLL};
183 use std::io::{Error, ErrorKind};
184
185 #[test]
186 fn drain_stops_on_persistent_error_instead_of_spinning() {
187 let mut calls = 0u32;
188 let mut out: Vec<u8> = Vec::new();
189 let capped = drain_accept(
190 || {
191 calls += 1;
192 Err::<u8, _>(Error::from(ErrorKind::Other))
194 },
195 &mut out,
196 MAX_ACCEPTS_PER_POLL,
197 );
198 assert_eq!(
199 calls, 1,
200 "a persistent accept error must break the drain, not loop forever"
201 );
202 assert!(out.is_empty());
203 assert!(!capped, "an errored drain leaves no known backlog");
204 }
205
206 #[test]
207 fn drain_collects_ready_connections_until_would_block() {
208 let mut seq = vec![
209 Ok(1u8),
210 Ok(2u8),
211 Err(Error::from(ErrorKind::WouldBlock)),
212 Ok(3u8),
213 ]
214 .into_iter();
215 let mut out = Vec::new();
216 let capped = drain_accept(|| seq.next().unwrap(), &mut out, MAX_ACCEPTS_PER_POLL);
217 assert_eq!(out, vec![1, 2]);
219 assert!(!capped, "a drained backlog must not be reported as pending");
220 }
221
222 #[test]
223 fn drain_stops_at_the_batch_limit() {
224 let mut accepted = 0usize;
226 let mut out = Vec::new();
227 let capped = drain_accept(
228 || {
229 accepted += 1;
230 Ok::<u8, Error>(0)
231 },
232 &mut out,
233 MAX_ACCEPTS_PER_POLL,
234 );
235 assert_eq!(out.len(), MAX_ACCEPTS_PER_POLL);
236 assert_eq!(
237 accepted, MAX_ACCEPTS_PER_POLL,
238 "the drain must not accept beyond the batch limit"
239 );
240 assert!(capped, "a capped drain must report the remaining backlog");
241 }
242}