shell_tunnel/pty/
async_adapter.rs1use std::io::{Read, Write};
8use tokio::sync::mpsc;
9use tracing::{debug, error, trace};
10
11pub struct AsyncPtyReader<R: Read + Send + 'static> {
15 reader: R,
16 tx: mpsc::Sender<Vec<u8>>,
17 buffer_size: usize,
18}
19
20impl<R: Read + Send + 'static> AsyncPtyReader<R> {
21 pub fn new(reader: R, tx: mpsc::Sender<Vec<u8>>) -> Self {
28 Self {
29 reader,
30 tx,
31 buffer_size: 4096,
32 }
33 }
34
35 pub fn with_buffer_size(mut self, size: usize) -> Self {
37 self.buffer_size = size;
38 self
39 }
40
41 pub async fn run(self) {
49 let buffer_size = self.buffer_size;
50 let mut reader = self.reader;
51 let tx = self.tx;
52
53 let result = tokio::task::spawn_blocking(move || {
54 let mut buf = vec![0u8; buffer_size];
55
56 loop {
57 match reader.read(&mut buf) {
58 Ok(0) => {
59 debug!("PTY reader: EOF");
60 break;
61 }
62 Ok(n) => {
63 trace!("PTY reader: read {} bytes", n);
64 if tx.blocking_send(buf[..n].to_vec()).is_err() {
65 debug!("PTY reader: channel closed");
66 break;
67 }
68 }
69 Err(e) => {
70 #[cfg(unix)]
72 if e.raw_os_error() == Some(libc::EIO) {
73 debug!("PTY reader: PTY closed (EIO)");
74 break;
75 }
76
77 if e.kind() == std::io::ErrorKind::BrokenPipe {
79 debug!("PTY reader: broken pipe");
80 break;
81 }
82
83 error!("PTY reader error: {}", e);
84 break;
85 }
86 }
87 }
88 })
89 .await;
90
91 if let Err(e) = result {
92 error!("PTY reader task panicked: {}", e);
93 }
94 }
95}
96
97pub struct AsyncPtyWriter<W: Write + Send + 'static> {
101 writer: W,
102 rx: mpsc::Receiver<Vec<u8>>,
103}
104
105impl<W: Write + Send + 'static> AsyncPtyWriter<W> {
106 pub fn new(writer: W, rx: mpsc::Receiver<Vec<u8>>) -> Self {
113 Self { writer, rx }
114 }
115
116 pub async fn run(self) {
123 let mut writer = self.writer;
124 let mut rx = self.rx;
125
126 let result = tokio::task::spawn_blocking(move || {
127 while let Some(data) = {
130 tokio::runtime::Handle::try_current()
133 .ok()
134 .and_then(|h| h.block_on(async { rx.recv().await }))
135 .or_else(|| {
136 rx.blocking_recv()
138 })
139 } {
140 trace!("PTY writer: writing {} bytes", data.len());
141 if let Err(e) = writer.write_all(&data) {
142 if e.kind() == std::io::ErrorKind::BrokenPipe {
143 debug!("PTY writer: broken pipe");
144 break;
145 }
146 error!("PTY writer error: {}", e);
147 break;
148 }
149 if let Err(e) = writer.flush() {
150 error!("PTY writer flush error: {}", e);
151 break;
152 }
153 }
154 debug!("PTY writer: channel closed");
155 })
156 .await;
157
158 if let Err(e) = result {
159 error!("PTY writer task panicked: {}", e);
160 }
161 }
162}
163
164#[cfg(test)]
165mod tests {
166 use super::*;
167 use std::io::Cursor;
168 use std::time::Duration;
169
170 #[tokio::test]
171 async fn test_async_reader_basic() {
172 let data = b"Hello, World!\nTest line 2\n";
174 let cursor = Cursor::new(data.to_vec());
175
176 let (tx, mut rx) = mpsc::channel(32);
177 let reader = AsyncPtyReader::new(cursor, tx);
178
179 let handle = tokio::spawn(reader.run());
181
182 let mut received = Vec::new();
184 while let Ok(Some(chunk)) =
185 tokio::time::timeout(Duration::from_millis(100), rx.recv()).await
186 {
187 received.extend(chunk);
188 }
189
190 let _ = tokio::time::timeout(Duration::from_millis(100), handle).await;
192
193 assert_eq!(received, data);
194 }
195
196 #[tokio::test]
197 async fn test_async_reader_empty() {
198 let cursor = Cursor::new(Vec::new());
199 let (tx, mut rx) = mpsc::channel(32);
200 let reader = AsyncPtyReader::new(cursor, tx);
201
202 let handle = tokio::spawn(reader.run());
203
204 let result = tokio::time::timeout(Duration::from_millis(100), rx.recv()).await;
206 assert!(result.is_ok());
207 assert!(result.unwrap().is_none()); let _ = handle.await;
210 }
211
212 #[tokio::test]
213 async fn test_async_writer_basic() {
214 let buffer = Vec::new();
215 let cursor = Cursor::new(buffer);
216
217 let (tx, rx) = mpsc::channel(32);
218 let writer = AsyncPtyWriter::new(cursor, rx);
219
220 tx.send(b"Hello".to_vec()).await.unwrap();
222 tx.send(b", World!".to_vec()).await.unwrap();
223 drop(tx); let handle = tokio::spawn(writer.run());
227 let _ = tokio::time::timeout(Duration::from_millis(500), handle).await;
228
229 }
232
233 #[tokio::test]
234 async fn test_channel_creation() {
235 let (reader_tx, mut reader_rx) = mpsc::channel::<Vec<u8>>(32);
237 let (writer_tx, _writer_rx) = mpsc::channel::<Vec<u8>>(32);
238
239 reader_tx.send(b"test".to_vec()).await.unwrap();
241 let received = reader_rx.recv().await.unwrap();
242 assert_eq!(received, b"test");
243
244 writer_tx.send(b"input".to_vec()).await.unwrap();
245 }
246
247 #[tokio::test]
248 async fn test_reader_channel_closed() {
249 let data = b"Some data that won't be fully read";
250 let cursor = Cursor::new(data.to_vec());
251
252 let (tx, rx) = mpsc::channel(1); let reader = AsyncPtyReader::new(cursor, tx);
254
255 drop(rx);
257
258 let handle = tokio::spawn(reader.run());
260 let result = tokio::time::timeout(Duration::from_millis(100), handle).await;
261 assert!(result.is_ok()); }
263}