1use 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
20pub 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 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 unsafe {
36 libc::fcntl(fd, libc::F_SETFL, libc::O_NONBLOCK);
37 libc::fcntl(fd, libc::F_SETFD, libc::FD_CLOEXEC);
38 }
39 }
40 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 pub fn wake(&self) {
51 let b = [1u8];
52 unsafe { libc::write(self.write.as_raw_fd(), b.as_ptr().cast(), 1) };
54 }
55
56 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 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
71pub 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
99pub trait Session {
101 fn outgoing(&mut self) -> Vec<String>;
103 fn incoming(&mut self, text: &str);
104 fn done(&self) -> bool {
106 false
107 }
108}
109
110pub const IDLE: Duration = Duration::from_secs(45);
113
114const 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
120pub 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
135pub 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
147pub 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 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 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 let ready = unsafe { libc::poll(fds.as_mut_ptr(), 1, 0) };
266 assert_eq!(ready, 0, "drained");
267 }
268}