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
13pub 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}