File size: 4,327 Bytes
c981f27
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
//! Shared output printer for synchronized writes to stdout/stderr.
//!
//! Prevents interleaving when multiple threads write to terminal output.

use std::io::{self, Stderr, Stdout, Write};
use std::sync::{Arc, Mutex};

use forge_domain::ConsoleWriter;

/// Thread-safe output printer that synchronizes writes to stdout/stderr.
///
/// Wraps writers in mutexes to prevent output interleaving when multiple
/// threads (e.g., streaming markdown and shell commands) write concurrently.
///
/// Generic over writer types `O` (stdout) and `E` (stderr) to support testing
/// with mock writers.
#[derive(Debug)]
pub struct StdConsoleWriter<O = Stdout, E = Stderr> {
    stdout: Arc<Mutex<O>>,
    stderr: Arc<Mutex<E>>,
}

impl<O, E> Clone for StdConsoleWriter<O, E> {
    fn clone(&self) -> Self {
        Self { stdout: self.stdout.clone(), stderr: self.stderr.clone() }
    }
}

impl Default for StdConsoleWriter<Stdout, Stderr> {
    fn default() -> Self {
        Self {
            stdout: Arc::new(Mutex::new(io::stdout())),
            stderr: Arc::new(Mutex::new(io::stderr())),
        }
    }
}

impl<O, E> StdConsoleWriter<O, E> {
    /// Creates a new OutputPrinter with custom writers.
    pub fn with_writers(stdout: O, stderr: E) -> Self {
        Self {
            stdout: Arc::new(Mutex::new(stdout)),
            stderr: Arc::new(Mutex::new(stderr)),
        }
    }
}

impl<O: Write + Send, E: Write + Send> ConsoleWriter for StdConsoleWriter<O, E> {
    fn write(&self, buf: &[u8]) -> io::Result<usize> {
        let mut guard = self.stdout.lock().unwrap_or_else(|e| e.into_inner());
        guard.write(buf)
    }

    fn write_err(&self, buf: &[u8]) -> io::Result<usize> {
        let mut guard = self.stderr.lock().unwrap_or_else(|e| e.into_inner());
        guard.write(buf)
    }

    fn flush(&self) -> io::Result<()> {
        let mut guard = self.stdout.lock().unwrap_or_else(|e| e.into_inner());
        guard.flush()
    }

    fn flush_err(&self) -> io::Result<()> {
        let mut guard = self.stderr.lock().unwrap_or_else(|e| e.into_inner());
        guard.flush()
    }
}

#[cfg(test)]
mod tests {
    use std::io::Cursor;
    use std::thread;

    use bstr::ByteSlice;

    use super::*;

    #[test]
    fn test_concurrent_writes_dont_interleave() {
        let stdout = Cursor::new(Vec::new());
        let stderr = Cursor::new(Vec::new());
        let printer = StdConsoleWriter::with_writers(stdout, stderr);
        let p1 = printer.clone();
        let p2 = printer.clone();

        let h1 = thread::spawn(move || {
            p1.write(b"AAAA").unwrap();
            p1.write(b"BBBB").unwrap();
            p1.flush().unwrap();
        });

        let h2 = thread::spawn(move || {
            p2.write(b"XXXX").unwrap();
            p2.write(b"ZZZZ").unwrap();
            p2.flush().unwrap();
        });

        h1.join().unwrap();
        h2.join().unwrap();

        // Verify output is one of the valid orderings where individual writes are
        // atomic but sequences can interleave. AAAA must come before BBBB, XXXX
        // must come before ZZZZ
        let actual = printer.stdout.lock().unwrap().get_ref().clone();
        let valid_orderings = [
            b"AAAABBBBXXXXZZZZ".to_vec(), // Thread 1 completes, then Thread 2
            b"XXXXZZZZAAAABBBB".to_vec(), // Thread 2 completes, then Thread 1
            b"AAAAXXXXBBBBZZZZ".to_vec(), // A, X, B, Z
            b"AAAAXXXXZZZZBBBB".to_vec(), // A, X, Z, B
            b"XXXXAAAABBBBZZZZ".to_vec(), // X, A, B, Z
            b"XXXXAAAAZZZZBBBB".to_vec(), // X, A, Z, B
        ];
        assert!(
            valid_orderings.contains(&actual),
            "Output was interleaved: {:?}",
            actual.as_slice().to_str_lossy()
        );
    }

    #[test]
    fn test_with_mock_writer() {
        let stdout = Cursor::new(Vec::new());
        let stderr = Cursor::new(Vec::new());
        let printer = StdConsoleWriter::with_writers(stdout, stderr);

        printer.write(b"hello").unwrap();
        printer.write_err(b"error").unwrap();

        let stdout_content = printer.stdout.lock().unwrap().get_ref().clone();
        let stderr_content = printer.stderr.lock().unwrap().get_ref().clone();

        assert_eq!(stdout_content, b"hello");
        assert_eq!(stderr_content, b"error");
    }
}