Skip to main content

can_motor_control/transport/
poller.rs

1//! Multi-bus poller built on `mio::Poll`.
2
3use std::os::fd::{BorrowedFd, RawFd};
4use std::time::Duration;
5
6use mio::unix::SourceFd;
7use mio::{Events, Interest, Poll, Token};
8
9use super::TransportError;
10
11/// Registers multiple bus fds and waits for any to become readable within a
12/// deadline. Buses with no `raw_fd` cannot be registered and must be polled
13/// out-of-band.
14pub struct BusPoller {
15    poll: Poll,
16    events: Events,
17}
18
19impl BusPoller {
20    /// Construct a poller with capacity for `n` simultaneous ready buses.
21    pub fn with_capacity(n: usize) -> Result<Self, TransportError> {
22        let poll = Poll::new().map_err(TransportError::Io)?;
23        Ok(Self {
24            poll,
25            events: Events::with_capacity(n.max(1)),
26        })
27    }
28
29    /// Register a fd under a token.
30    pub fn register(&self, token: Token, fd: RawFd) -> Result<(), TransportError> {
31        // SAFETY: BorrowedFd is constructed for the duration of register;
32        // the caller is responsible for keeping the fd alive while registered.
33        let mut src = SourceFd(&fd);
34        self.poll
35            .registry()
36            .register(&mut src, token, Interest::READABLE)
37            .map_err(TransportError::Io)?;
38        // BorrowedFd is only used for its lifetime guarantee; we drop it now.
39        let _ = unsafe { BorrowedFd::borrow_raw(fd) };
40        Ok(())
41    }
42
43    /// Deregister a fd; call when the bus is going away.
44    pub fn deregister(&self, fd: RawFd) -> Result<(), TransportError> {
45        let mut src = SourceFd(&fd);
46        self.poll
47            .registry()
48            .deregister(&mut src)
49            .map_err(TransportError::Io)
50    }
51
52    /// Block up to `deadline` waiting for any registered fd to become readable.
53    /// Returns the tokens of the ready fds in arrival order.
54    pub fn wait(&mut self, deadline: Duration) -> Result<Vec<Token>, TransportError> {
55        match self.poll.poll(&mut self.events, Some(deadline)) {
56            Ok(()) => {}
57            Err(e) if e.kind() == std::io::ErrorKind::Interrupted => {
58                // EINTR: treat as a spurious wake-up and return empty.
59                return Ok(Vec::new());
60            }
61            Err(e) => return Err(TransportError::Io(e)),
62        }
63        Ok(self.events.iter().map(|ev| ev.token()).collect())
64    }
65}
66
67#[cfg(test)]
68mod tests {
69    use super::*;
70    use std::os::fd::AsRawFd;
71    use std::time::Instant;
72
73    fn make_pipe() -> (std::fs::File, std::fs::File) {
74        // Use a simple unix pipe for fd-readable tests.
75        let mut pipefd = [0i32; 2];
76        // SAFETY: standard pipe() with a valid output array.
77        let rc = unsafe { libc::pipe(pipefd.as_mut_ptr()) };
78        assert!(rc == 0);
79        unsafe {
80            (
81                std::fs::File::from_raw_fd_unchecked(pipefd[0]),
82                std::fs::File::from_raw_fd_unchecked(pipefd[1]),
83            )
84        }
85    }
86
87    trait FromRawFdUnchecked {
88        unsafe fn from_raw_fd_unchecked(fd: i32) -> Self;
89    }
90    impl FromRawFdUnchecked for std::fs::File {
91        unsafe fn from_raw_fd_unchecked(fd: i32) -> Self {
92            use std::os::fd::FromRawFd;
93            std::fs::File::from_raw_fd(fd)
94        }
95    }
96
97    #[test]
98    fn quiet_buses_deadline_expires() {
99        let mut p = BusPoller::with_capacity(4).unwrap();
100        let (r, _w) = make_pipe();
101        p.register(Token(0), r.as_raw_fd()).unwrap();
102        let t0 = Instant::now();
103        let tokens = p.wait(Duration::from_millis(5)).unwrap();
104        let elapsed = t0.elapsed();
105        assert!(tokens.is_empty(), "tokens: {tokens:?}");
106        assert!(
107            elapsed < Duration::from_millis(50),
108            "took too long: {elapsed:?}"
109        );
110    }
111
112    #[test]
113    fn wake_on_readable_pipe() {
114        use std::io::Write;
115        let mut p = BusPoller::with_capacity(4).unwrap();
116        let (r, mut w) = make_pipe();
117        p.register(Token(7), r.as_raw_fd()).unwrap();
118        w.write_all(b"x").unwrap();
119        let tokens = p.wait(Duration::from_millis(100)).unwrap();
120        assert!(tokens.contains(&Token(7)), "tokens: {tokens:?}");
121    }
122}