Skip to main content

koan_core/remote/
wire.rs

1//! WebSockets that sleep until there is something to do.
2//!
3//! A link carries two kinds of traffic: what arrives from the other end, and
4//! what this end has to say because something here moved — the player, the
5//! queue, a command the user gave. The socket's descriptor answers the first;
6//! a [`Waker`] (a pipe) answers the second. One `poll` on both, and a quiet
7//! link schedules nothing between its keep-alive pings.
8
9use std::io::{Read, Write};
10use std::net::TcpStream;
11use std::os::fd::{AsRawFd, FromRawFd, OwnedFd, RawFd};
12use std::sync::atomic::{AtomicU64, Ordering};
13use std::sync::{Arc, Weak};
14use std::time::{Duration, Instant};
15
16use parking_lot::Mutex;
17use tungstenite::stream::MaybeTlsStream;
18use tungstenite::{Message, WebSocket};
19
20/// A pipe a thread can be woken through while it waits on a socket.
21pub struct Waker {
22    read: OwnedFd,
23    write: OwnedFd,
24}
25
26impl Waker {
27    pub fn new() -> std::io::Result<Arc<Self>> {
28        let mut fds = [0; 2];
29        // SAFETY: `fds` has room for the two descriptors pipe writes.
30        if unsafe { libc::pipe(fds.as_mut_ptr()) } != 0 {
31            return Err(std::io::Error::last_os_error());
32        }
33        for fd in fds {
34            // SAFETY: both are descriptors pipe just returned.
35            unsafe {
36                libc::fcntl(fd, libc::F_SETFL, libc::O_NONBLOCK);
37                libc::fcntl(fd, libc::F_SETFD, libc::FD_CLOEXEC);
38            }
39        }
40        // SAFETY: each descriptor is owned by exactly one OwnedFd from here.
41        Ok(Arc::new(unsafe {
42            Self {
43                read: OwnedFd::from_raw_fd(fds[0]),
44                write: OwnedFd::from_raw_fd(fds[1]),
45            }
46        }))
47    }
48
49    /// Wake whoever waits. A full pipe already says so, so it is not an error.
50    pub fn wake(&self) {
51        let b = [1u8];
52        // SAFETY: writes one byte from a live buffer to a descriptor we own.
53        unsafe { libc::write(self.write.as_raw_fd(), b.as_ptr().cast(), 1) };
54    }
55
56    /// The end to wait on for a wake.
57    pub fn read_fd(&self) -> RawFd {
58        self.read.as_raw_fd()
59    }
60
61    pub fn drain(&self) {
62        let mut buf = [0u8; 64];
63        // SAFETY: reads into a live buffer of the length given.
64        while unsafe { libc::read(self.read.as_raw_fd(), buf.as_mut_ptr().cast(), buf.len()) } > 0 {
65        }
66    }
67}
68
69static ON_ENGINE: Mutex<Vec<Weak<Waker>>> = Mutex::new(Vec::new());
70
71/// Wake `waker` whenever the engine says something moved, for as long as it
72/// lives. One thread serves every waker in the process.
73pub fn wake_on_engine_change(waker: &Arc<Waker>) {
74    let mut list = ON_ENGINE.lock();
75    let first = list.is_empty();
76    list.push(Arc::downgrade(waker));
77    drop(list);
78    if first {
79        std::thread::Builder::new()
80            .name("koan-wire".into())
81            .spawn(|| {
82                let signal = crate::signal::engine_changed();
83                let mut seen = signal.generation();
84                loop {
85                    seen = signal.wait(seen);
86                    ON_ENGINE.lock().retain(|w| match w.upgrade() {
87                        Some(w) => {
88                            w.wake();
89                            true
90                        }
91                        None => false,
92                    });
93                }
94            })
95            .expect("failed to spawn the wire thread");
96    }
97}
98
99/// One end of a conversation over a socket.
100pub trait Session {
101    /// What to send now. Asked on every wake and after every message read.
102    fn outgoing(&mut self) -> Vec<String>;
103    fn incoming(&mut self, text: &str);
104    /// Whether to hang up. Asked on every wake.
105    fn done(&self) -> bool {
106        false
107    }
108}
109
110/// A link that has heard nothing for this long pings, so a dead connection is
111/// noticed rather than waited on forever.
112pub const IDLE: Duration = Duration::from_secs(45);
113
114/// How long a connection has to answer a probe.
115const PROBE_WAIT: Duration = Duration::from_secs(2);
116
117static PROBES: AtomicU64 = AtomicU64::new(0);
118static SESSIONS: Mutex<Vec<Weak<Waker>>> = Mutex::new(Vec::new());
119
120/// Have every connection prove it is alive now, and drop those that do not
121/// answer within two seconds. For an app coming back from suspension: its
122/// sockets look open, but the far end gave up on them long ago, and waiting
123/// for the idle ping to find out keeps the other devices out of sight.
124pub fn probe_all() {
125    PROBES.fetch_add(1, Ordering::Relaxed);
126    SESSIONS.lock().retain(|w| match w.upgrade() {
127        Some(w) => {
128            w.wake();
129            true
130        }
131        None => false,
132    });
133}
134
135/// The descriptor under a client's socket, and put it in non-blocking mode.
136pub fn prepare(stream: &MaybeTlsStream<TcpStream>) -> Result<RawFd, String> {
137    let tcp = match stream {
138        MaybeTlsStream::Plain(s) => s,
139        MaybeTlsStream::Rustls(s) => s.get_ref(),
140        _ => return Err("unsupported stream".into()),
141    };
142    tcp.set_nonblocking(true).map_err(|e| e.to_string())?;
143    let _ = tcp.set_nodelay(true);
144    Ok(tcp.as_raw_fd())
145}
146
147/// Run `session` over `socket` until either end goes. `fd` is the socket's
148/// descriptor, already non-blocking; `waker` interrupts the wait.
149pub fn drive<S: Read + Write>(
150    socket: &mut WebSocket<S>,
151    fd: RawFd,
152    waker: &Arc<Waker>,
153    session: &mut impl Session,
154) -> Result<(), String> {
155    let mut heard = Instant::now();
156    let mut pinged = false;
157    let mut probes = PROBES.load(Ordering::Relaxed);
158    let mut probed: Option<Instant> = None;
159    {
160        let mut sessions = SESSIONS.lock();
161        sessions.retain(|w| w.strong_count() > 0);
162        sessions.push(Arc::downgrade(waker));
163    }
164    loop {
165        if session.done() {
166            let _ = socket.close(None);
167            let _ = socket.flush();
168            return Ok(());
169        }
170        for text in session.outgoing() {
171            match socket.write(Message::Text(text.into())) {
172                Ok(()) => {}
173                Err(tungstenite::Error::Io(e)) if e.kind() == std::io::ErrorKind::WouldBlock => {}
174                Err(e) => return Err(e.to_string()),
175            }
176        }
177        let mut read_any = false;
178        loop {
179            match socket.read() {
180                Ok(Message::Text(text)) => {
181                    (heard, pinged, read_any) = (Instant::now(), false, true);
182                    session.incoming(&text);
183                }
184                Ok(Message::Close(_)) => return Err("closed by the other end".into()),
185                Ok(_) => (heard, pinged) = (Instant::now(), false),
186                Err(tungstenite::Error::Io(e)) if e.kind() == std::io::ErrorKind::WouldBlock => {
187                    break;
188                }
189                Err(e) => return Err(e.to_string()),
190            }
191        }
192        // A message read may have something to answer at once.
193        if read_any {
194            continue;
195        }
196        let blocked = match socket.flush() {
197            Ok(()) => false,
198            Err(tungstenite::Error::Io(e)) if e.kind() == std::io::ErrorKind::WouldBlock => true,
199            Err(e) => return Err(e.to_string()),
200        };
201        let now = PROBES.load(Ordering::Relaxed);
202        if now != probes {
203            probes = now;
204            let _ = socket.write(Message::Ping(Vec::new().into()));
205            let _ = socket.flush();
206            probed = Some(Instant::now());
207            continue;
208        }
209        let mut left = IDLE.saturating_sub(heard.elapsed());
210        if let Some(at) = probed {
211            if heard >= at {
212                probed = None;
213            } else {
214                let wait = PROBE_WAIT.saturating_sub(at.elapsed());
215                if wait.is_zero() {
216                    return Err("no answer after waking".into());
217                }
218                left = left.min(wait);
219            }
220        }
221        if left.is_zero() {
222            if pinged {
223                return Err("no answer to a ping".into());
224            }
225            let _ = socket.write(Message::Ping(Vec::new().into()));
226            (heard, pinged) = (Instant::now(), true);
227            continue;
228        }
229        let mut fds = [
230            libc::pollfd {
231                fd,
232                events: libc::POLLIN | if blocked { libc::POLLOUT } else { 0 },
233                revents: 0,
234            },
235            libc::pollfd {
236                fd: waker.read.as_raw_fd(),
237                events: libc::POLLIN,
238                revents: 0,
239            },
240        ];
241        let ms = left.as_millis().min(i32::MAX as u128) as i32;
242        // SAFETY: `fds` is a live array of the length given.
243        unsafe { libc::poll(fds.as_mut_ptr(), fds.len() as _, ms) };
244        waker.drain();
245    }
246}
247
248#[cfg(test)]
249mod tests {
250    use super::*;
251
252    #[test]
253    fn a_waker_can_be_woken_more_than_it_is_read() {
254        let waker = Waker::new().unwrap();
255        for _ in 0..100_000 {
256            waker.wake();
257        }
258        waker.drain();
259        let mut fds = [libc::pollfd {
260            fd: waker.read.as_raw_fd(),
261            events: libc::POLLIN,
262            revents: 0,
263        }];
264        // SAFETY: as above.
265        let ready = unsafe { libc::poll(fds.as_mut_ptr(), 1, 0) };
266        assert_eq!(ready, 0, "drained");
267    }
268}