Skip to main content

muxr_client/
stdout_worker.rs

1use std::collections::VecDeque;
2use std::io::Write;
3use std::sync::Arc;
4use std::sync::Condvar;
5use std::sync::Mutex;
6use std::sync::MutexGuard;
7use std::thread;
8
9use rootcause::report;
10
11const QUEUED_TRANSACTION_BYTE_LIMIT: usize = 4 * 1024 * 1024;
12
13/// A single stdout owner that reports each successful render flush.
14pub struct StdoutWorker {
15    shared: Arc<Shared>,
16    handle: Option<thread::JoinHandle<()>>,
17}
18
19#[derive(Clone)]
20pub struct StdoutSender {
21    shared: Arc<Shared>,
22}
23
24struct Shared {
25    completed: tokio::sync::mpsc::UnboundedSender<()>,
26    state: Mutex<State>,
27    wake: Condvar,
28}
29
30struct State {
31    closed: bool,
32    failed: Option<String>,
33    output: VecDeque<OutputCmd>,
34    queued_bytes: usize,
35}
36
37enum OutputCmd {
38    Render(Vec<u8>),
39}
40
41impl StdoutWorker {
42    pub fn spawn() -> (
43        StdoutSender,
44        Self,
45        tokio::sync::oneshot::Receiver<String>,
46        tokio::sync::mpsc::UnboundedReceiver<()>,
47    ) {
48        let (failure_sender, failure_receiver) = tokio::sync::oneshot::channel();
49        let (completed_sender, completed_receiver) = tokio::sync::mpsc::unbounded_channel();
50        let shared = Arc::new(Shared {
51            completed: completed_sender,
52            state: Mutex::new(State {
53                closed: false,
54                failed: None,
55                output: VecDeque::new(),
56                queued_bytes: 0,
57            }),
58            wake: Condvar::new(),
59        });
60        let worker_shared = Arc::clone(&shared);
61        let handle = thread::spawn(move || {
62            let mut stdout = std::io::stdout();
63            run(&worker_shared, &mut stdout, failure_sender);
64        });
65        (
66            StdoutSender {
67                shared: Arc::clone(&shared),
68            },
69            Self {
70                shared,
71                handle: Some(handle),
72            },
73            failure_receiver,
74            completed_receiver,
75        )
76    }
77}
78
79impl Drop for StdoutWorker {
80    fn drop(&mut self) {
81        if let Ok(mut state) = self::lock_state(&self.shared) {
82            state.closed = true;
83            drop(state);
84            self.shared.wake.notify_all();
85        }
86        if let Some(handle) = self.handle.take() {
87            let _joined = handle.join();
88        }
89    }
90}
91
92impl StdoutSender {
93    pub fn send_render(&self, transaction: Vec<u8>) -> rootcause::Result<()> {
94        if transaction.is_empty() {
95            return Ok(());
96        }
97        let mut state = self::lock_state(&self.shared)?;
98        ensure_open(&state)?;
99        self::reserve_queued_bytes(&state, transaction.len())?;
100        state.queued_bytes = state
101            .queued_bytes
102            .checked_add(transaction.len())
103            .ok_or_else(|| report!("muxr client stdout queued byte count overflowed"))?;
104        state.output.push_back(OutputCmd::Render(transaction));
105        drop(state);
106        self.shared.wake.notify_one();
107        Ok(())
108    }
109}
110
111fn reserve_queued_bytes(state: &State, additional_bytes: usize) -> rootcause::Result<()> {
112    let next = state
113        .queued_bytes
114        .checked_add(additional_bytes)
115        .ok_or_else(|| report!("muxr client stdout queued byte count overflowed"))?;
116    if next > QUEUED_TRANSACTION_BYTE_LIMIT {
117        return Err(report!("muxr client stdout queued byte budget is exhausted"));
118    }
119    Ok(())
120}
121
122impl OutputCmd {
123    const fn len(&self) -> usize {
124        match self {
125            Self::Render(transaction) => transaction.len(),
126        }
127    }
128}
129
130fn ensure_open(state: &State) -> rootcause::Result<()> {
131    if let Some(error) = &state.failed {
132        return Err(report!("muxr client stdout worker failed").attach(error.clone()));
133    }
134    if state.closed {
135        return Err(report!("muxr client stdout worker is closed"));
136    }
137    Ok(())
138}
139
140fn lock_state(shared: &Shared) -> rootcause::Result<MutexGuard<'_, State>> {
141    shared
142        .state
143        .lock()
144        .map_err(|_| report!("muxr stdout worker state poisoned"))
145}
146
147fn run(shared: &Shared, stdout: &mut impl Write, failure_sender: tokio::sync::oneshot::Sender<String>) {
148    loop {
149        let cmd = {
150            let Ok(mut state) = self::lock_state(shared) else {
151                return;
152            };
153            while !state.closed && state.output.is_empty() {
154                let Ok(next_state) = shared.wake.wait(state) else {
155                    return;
156                };
157                state = next_state;
158            }
159            if state.closed && state.output.is_empty() {
160                return;
161            }
162            let cmd = state.output.pop_front();
163            if let Some(cmd) = cmd.as_ref() {
164                state.queued_bytes = state.queued_bytes.saturating_sub(cmd.len());
165            }
166            cmd
167        };
168        let Some(cmd) = cmd else {
169            continue;
170        };
171        let OutputCmd::Render(transaction) = cmd;
172        if let Err(error) = stdout.write_all(&transaction).and_then(|()| stdout.flush()) {
173            if let Ok(mut state) = self::lock_state(shared) {
174                state.failed = Some(error.to_string());
175                state.closed = true;
176                drop(state);
177            }
178            shared.wake.notify_all();
179            let _sent = failure_sender.send(error.to_string());
180            return;
181        }
182        let _sent = shared.completed.send(());
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use test_that::prelude::*;
189
190    use super::*;
191
192    #[test]
193    fn test_queued_byte_budget_rejects_another_transaction() {
194        let state = State {
195            closed: false,
196            failed: None,
197            output: VecDeque::new(),
198            queued_bytes: QUEUED_TRANSACTION_BYTE_LIMIT,
199        };
200
201        assert_that!(reserve_queued_bytes(&state, 1).is_err(), eq(true));
202    }
203
204    #[test]
205    fn test_run_when_render_flushes_sends_completion_after_output() {
206        let (completed, mut completed_receiver) = tokio::sync::mpsc::unbounded_channel();
207        let shared = Shared {
208            completed,
209            state: Mutex::new(State {
210                closed: true,
211                failed: None,
212                output: VecDeque::from([OutputCmd::Render(b"render".to_vec())]),
213                queued_bytes: b"render".len(),
214            }),
215            wake: Condvar::new(),
216        };
217        let (failure_sender, mut failure_receiver) = tokio::sync::oneshot::channel();
218        let mut output = Vec::new();
219
220        run(&shared, &mut output, failure_sender);
221
222        assert_that!(output, eq(b"render".to_vec()));
223        assert_that!(completed_receiver.try_recv(), eq(Ok(())));
224        assert_that!(failure_receiver.try_recv().is_err(), eq(true));
225    }
226
227    #[test]
228    fn test_run_when_render_flush_blocks_completion_until_flush_finishes() -> rootcause::Result<()> {
229        let (completed, mut completed_receiver) = tokio::sync::mpsc::unbounded_channel();
230        let shared = Arc::new(Shared {
231            completed,
232            state: Mutex::new(State {
233                closed: true,
234                failed: None,
235                output: VecDeque::from([OutputCmd::Render(b"render".to_vec())]),
236                queued_bytes: b"render".len(),
237            }),
238            wake: Condvar::new(),
239        });
240        let (flush_started_sender, flush_started_receiver) = std::sync::mpsc::channel();
241        let (flush_release_sender, flush_release_receiver) = std::sync::mpsc::channel();
242        let (failure_sender, _failure_receiver) = tokio::sync::oneshot::channel();
243        let worker_shared = Arc::clone(&shared);
244        let handle = thread::spawn(move || {
245            let mut output = BlockingWriter {
246                flush_release_receiver,
247                flush_started_sender,
248            };
249            run(&worker_shared, &mut output, failure_sender);
250        });
251
252        flush_started_receiver.recv()?;
253        assert_that!(
254            completed_receiver.try_recv(),
255            err(matches_pattern!(tokio::sync::mpsc::error::TryRecvError::Empty))
256        );
257        flush_release_sender.send(())?;
258        handle
259            .join()
260            .map_err(|_| report!("muxr stdout blocking-writer test thread panicked"))?;
261
262        assert_that!(completed_receiver.try_recv(), eq(Ok(())));
263        Ok(())
264    }
265
266    #[test]
267    fn test_run_when_render_write_fails_marks_worker_failed_without_completion() -> rootcause::Result<()> {
268        let (completed, mut completed_receiver) = tokio::sync::mpsc::unbounded_channel();
269        let shared = Shared {
270            completed,
271            state: Mutex::new(State {
272                closed: true,
273                failed: None,
274                output: VecDeque::from([OutputCmd::Render(b"render".to_vec())]),
275                queued_bytes: b"render".len(),
276            }),
277            wake: Condvar::new(),
278        };
279        let (failure_sender, mut failure_receiver) = tokio::sync::oneshot::channel();
280
281        run(&shared, &mut FailingWriter, failure_sender);
282
283        assert_that!(failure_receiver.try_recv().is_ok(), eq(true));
284        assert_that!(
285            completed_receiver.try_recv(),
286            err(matches_pattern!(tokio::sync::mpsc::error::TryRecvError::Empty))
287        );
288        let state = lock_state(&shared)?;
289        assert_that!(state.failed.is_some(), eq(true));
290        assert_that!(state.closed, eq(true));
291        drop(state);
292        Ok(())
293    }
294
295    struct BlockingWriter {
296        flush_release_receiver: std::sync::mpsc::Receiver<()>,
297        flush_started_sender: std::sync::mpsc::Sender<()>,
298    }
299
300    impl Write for BlockingWriter {
301        fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
302            Ok(buffer.len())
303        }
304
305        fn flush(&mut self) -> std::io::Result<()> {
306            self.flush_started_sender
307                .send(())
308                .map_err(|_| std::io::Error::other("muxr stdout flush observer disconnected"))?;
309            self.flush_release_receiver
310                .recv()
311                .map_err(|_| std::io::Error::other("muxr stdout flush release disconnected"))
312        }
313    }
314
315    struct FailingWriter;
316
317    impl Write for FailingWriter {
318        fn write(&mut self, _buffer: &[u8]) -> std::io::Result<usize> {
319            Err(std::io::Error::other("injected muxr stdout write failure"))
320        }
321
322        fn flush(&mut self) -> std::io::Result<()> {
323            Ok(())
324        }
325    }
326}