1use crate::security::sanitize_text;
7use std::io::{self, BufRead, Write};
8use std::sync::atomic::{AtomicBool, Ordering};
9use std::sync::Arc;
10
11enum ReadLineOutcome {
12 Line,
13 Eof,
14}
15
16#[derive(Debug, Clone)]
18pub struct StdioConfig {
19 pub max_message_size: usize,
21 pub auto_flush: bool,
23 pub stderr_logging: bool,
25 pub buffer_size: usize,
27}
28
29impl Default for StdioConfig {
30 fn default() -> Self {
31 Self {
32 max_message_size: 10 * 1024 * 1024, auto_flush: true,
34 stderr_logging: true,
35 buffer_size: 4096,
36 }
37 }
38}
39
40pub struct StdioTransport {
42 read_buffer: String,
44 config: StdioConfig,
46 shutdown: Arc<AtomicBool>,
48}
49
50impl Default for StdioTransport {
51 fn default() -> Self {
52 Self::new()
53 }
54}
55
56impl StdioTransport {
57 pub fn new() -> Self {
59 Self::with_config(StdioConfig::default())
60 }
61
62 pub fn with_config(config: StdioConfig) -> Self {
64 Self {
65 read_buffer: String::with_capacity(config.buffer_size),
66 config,
67 shutdown: Arc::new(AtomicBool::new(false)),
68 }
69 }
70
71 pub fn with_auto_flush(mut self, auto_flush: bool) -> Self {
73 self.config.auto_flush = auto_flush;
74 self
75 }
76
77 pub fn shutdown_handle(&self) -> Arc<AtomicBool> {
79 Arc::clone(&self.shutdown)
80 }
81
82 pub fn shutdown(&self) {
84 self.shutdown.store(true, Ordering::SeqCst);
85 }
86
87 pub fn is_shutdown(&self) -> bool {
89 self.shutdown.load(Ordering::SeqCst)
90 }
91
92 pub fn read_message<R: BufRead>(&mut self, reader: &mut R) -> io::Result<Option<String>> {
94 loop {
95 if self.is_shutdown() {
96 return Ok(None);
97 }
98
99 self.read_buffer.clear();
100
101 match self.read_line_bounded(reader)? {
102 ReadLineOutcome::Eof => {
103 self.log_stderr("EOF received on stdin, initiating shutdown");
105 self.shutdown();
106 return Ok(None);
107 }
108 ReadLineOutcome::Line => {}
109 }
110
111 let message = self.read_buffer.trim_end().to_string();
113 if message.is_empty() {
114 continue;
115 }
116
117 return Ok(Some(message));
118 }
119 }
120
121 fn read_line_bounded<R: BufRead>(&mut self, reader: &mut R) -> io::Result<ReadLineOutcome> {
122 loop {
123 let available = reader.fill_buf()?;
124 if available.is_empty() {
125 if self.read_buffer.is_empty() {
126 return Ok(ReadLineOutcome::Eof);
127 }
128 return Ok(ReadLineOutcome::Line);
129 }
130
131 let take = available
132 .iter()
133 .position(|byte| *byte == b'\n')
134 .map(|index| index + 1)
135 .unwrap_or(available.len());
136
137 let next_len = self.read_buffer.len().saturating_add(take);
138 if next_len > self.config.max_message_size {
139 let remaining = self
140 .config
141 .max_message_size
142 .saturating_sub(self.read_buffer.len());
143 let allowed = remaining.min(available.len());
144 let consume_len = (allowed + 1).min(available.len());
145 if allowed > 0 {
146 let text = std::str::from_utf8(&available[..allowed]).map_err(|_| {
147 io::Error::new(io::ErrorKind::InvalidData, "Invalid UTF-8 in stdin message")
148 })?;
149 self.read_buffer.push_str(text);
150 }
151 reader.consume(consume_len);
152 self.log_error("stdin message exceeded maximum size");
153 self.shutdown();
154 self.read_buffer.clear();
155 return Err(io::Error::new(
156 io::ErrorKind::InvalidData,
157 "Message exceeds maximum size",
158 ));
159 }
160
161 let text = std::str::from_utf8(&available[..take]).map_err(|_| {
162 io::Error::new(io::ErrorKind::InvalidData, "Invalid UTF-8 in stdin message")
163 })?;
164 self.read_buffer.push_str(text);
165 reader.consume(take);
166
167 if self.read_buffer.ends_with('\n') {
168 return Ok(ReadLineOutcome::Line);
169 }
170 }
171 }
172
173 pub fn write_message<W: Write>(&self, writer: &mut W, message: &str) -> io::Result<()> {
175 if self.is_shutdown() {
176 return Err(io::Error::new(
177 io::ErrorKind::BrokenPipe,
178 "transport is shut down",
179 ));
180 }
181 if message.len() > self.config.max_message_size {
182 self.log_error("stdout message exceeded maximum size");
183 return Err(io::Error::new(
184 io::ErrorKind::InvalidData,
185 "Message exceeds maximum size",
186 ));
187 }
188 writeln!(writer, "{}", message)?;
189 if self.config.auto_flush {
190 writer.flush()?;
191 }
192 Ok(())
193 }
194
195 pub fn read_all_messages<R: BufRead>(&mut self, reader: &mut R) -> io::Result<Vec<String>> {
197 let mut messages = Vec::new();
198 while !self.is_shutdown() {
199 match self.read_message(reader)? {
200 Some(msg) => messages.push(msg),
201 None => break,
202 }
203 }
204 Ok(messages)
205 }
206
207 pub fn log_stderr(&self, message: &str) {
209 if self.config.stderr_logging {
210 let _ = writeln!(io::stderr(), "{}", Self::format_log_line("DCP", message));
211 }
212 }
213
214 pub fn log_error(&self, message: &str) {
216 if self.config.stderr_logging {
217 let _ = writeln!(
218 io::stderr(),
219 "{}",
220 Self::format_log_line("DCP ERROR", message)
221 );
222 }
223 }
224
225 pub fn log_debug(&self, message: &str) {
227 if self.config.stderr_logging {
228 let _ = writeln!(
229 io::stderr(),
230 "{}",
231 Self::format_log_line("DCP DEBUG", message)
232 );
233 }
234 }
235
236 pub fn format_log_line(prefix: &str, message: &str) -> String {
238 format!("[{}] {}", prefix, sanitize_text(message))
239 }
240}
241
242pub struct MessageFramer {
244 buffer: Vec<u8>,
246 max_size: usize,
248}
249
250impl Default for MessageFramer {
251 fn default() -> Self {
252 Self::new()
253 }
254}
255
256impl MessageFramer {
257 pub fn new() -> Self {
259 Self {
260 buffer: Vec::with_capacity(4096),
261 max_size: 10 * 1024 * 1024, }
263 }
264
265 pub fn with_max_size(mut self, max_size: usize) -> Self {
267 self.max_size = max_size;
268 self
269 }
270
271 pub fn feed(&mut self, data: &[u8]) -> io::Result<Vec<String>> {
273 let mut messages = Vec::new();
274
275 for &byte in data {
276 if byte == b'\n' {
277 if !self.buffer.is_empty() {
279 let message = String::from_utf8(std::mem::take(&mut self.buffer))
280 .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
281 let trimmed = message.trim().to_string();
282 if !trimmed.is_empty() {
283 messages.push(trimmed);
284 }
285 }
286 } else {
287 if self.buffer.len() >= self.max_size {
289 self.buffer.clear();
290 return Err(io::Error::new(
291 io::ErrorKind::InvalidData,
292 "Message exceeds maximum size",
293 ));
294 }
295 self.buffer.push(byte);
296 }
297 }
298
299 Ok(messages)
300 }
301
302 pub fn has_partial(&self) -> bool {
304 !self.buffer.is_empty()
305 }
306
307 pub fn buffer_size(&self) -> usize {
309 self.buffer.len()
310 }
311
312 pub fn clear(&mut self) {
314 self.buffer.clear();
315 }
316}
317
318pub fn frame_message(message: &str) -> String {
320 format!("{}\n", message)
321}
322
323pub fn unframe_message(data: &str) -> &str {
325 data.trim_end_matches('\n').trim_end_matches('\r')
326}
327
328#[cfg(feature = "async-stdio")]
330pub mod async_transport {
331 use super::*;
332 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
333 use tokio::sync::mpsc;
334
335 pub struct AsyncStdioTransport {
337 config: StdioConfig,
338 shutdown: Arc<AtomicBool>,
339 }
340
341 impl AsyncStdioTransport {
342 pub fn new(config: StdioConfig) -> Self {
344 Self {
345 config,
346 shutdown: Arc::new(AtomicBool::new(false)),
347 }
348 }
349
350 pub fn shutdown_handle(&self) -> Arc<AtomicBool> {
352 Arc::clone(&self.shutdown)
353 }
354
355 pub async fn run(self) -> io::Result<(mpsc::Receiver<String>, mpsc::Sender<String>)> {
357 let (in_tx, in_rx) = mpsc::channel(100);
358 let (out_tx, mut out_rx) = mpsc::channel::<String>(100);
359
360 let shutdown = Arc::clone(&self.shutdown);
361 let config = self.config.clone();
362
363 tokio::spawn(async move {
365 let stdin = tokio::io::stdin();
366 let mut reader = BufReader::new(stdin);
367 let mut line = String::new();
368
369 loop {
370 if shutdown.load(Ordering::SeqCst) {
371 break;
372 }
373
374 line.clear();
375 match reader.read_line(&mut line).await {
376 Ok(0) => {
377 if config.stderr_logging {
379 eprintln!("[DCP] EOF on stdin");
380 }
381 break;
382 }
383 Ok(_) => {
384 let msg = line.trim().to_string();
385 if !msg.is_empty() {
386 if in_tx.send(msg).await.is_err() {
387 break;
388 }
389 }
390 }
391 Err(e) => {
392 if config.stderr_logging {
393 eprintln!("[DCP ERROR] stdin read error: {}", e);
394 }
395 break;
396 }
397 }
398 }
399 });
400
401 tokio::spawn(async move {
403 let mut stdout = tokio::io::stdout();
404
405 while let Some(msg) = out_rx.recv().await {
406 let framed = format!("{}\n", msg);
407 if let Err(e) = stdout.write_all(framed.as_bytes()).await {
408 eprintln!("[DCP ERROR] stdout write error: {}", e);
409 break;
410 }
411 if let Err(e) = stdout.flush().await {
412 eprintln!("[DCP ERROR] stdout flush error: {}", e);
413 break;
414 }
415 }
416 });
417
418 Ok((in_rx, out_tx))
419 }
420 }
421}
422
423#[cfg(test)]
424mod tests {
425 use super::*;
426 use std::io::Cursor;
427
428 #[test]
429 fn test_stdio_transport_read_message() {
430 let mut transport = StdioTransport::new();
431 let input = r#"{"jsonrpc":"2.0","method":"test","id":1}
432"#;
433 let mut reader = Cursor::new(input);
434
435 let message = transport.read_message(&mut reader).unwrap();
436 assert_eq!(
437 message,
438 Some(r#"{"jsonrpc":"2.0","method":"test","id":1}"#.to_string())
439 );
440 }
441
442 #[test]
443 fn test_stdio_transport_read_eof() {
444 let mut transport = StdioTransport::new();
445 let mut reader = Cursor::new("");
446
447 let message = transport.read_message(&mut reader).unwrap();
448 assert_eq!(message, None);
449 assert!(transport.is_shutdown()); }
451
452 #[test]
453 fn test_stdio_transport_write_message() {
454 let transport = StdioTransport::new();
455 let mut output = Vec::new();
456
457 transport
458 .write_message(&mut output, r#"{"jsonrpc":"2.0","result":{},"id":1}"#)
459 .unwrap();
460
461 assert_eq!(
462 String::from_utf8(output).unwrap(),
463 "{\"jsonrpc\":\"2.0\",\"result\":{},\"id\":1}\n"
464 );
465 }
466
467 #[test]
468 fn test_stdio_transport_read_multiple() {
469 let mut transport = StdioTransport::new();
470 let input = r#"{"id":1}
471{"id":2}
472{"id":3}
473"#;
474 let mut reader = Cursor::new(input);
475
476 let messages = transport.read_all_messages(&mut reader).unwrap();
477 assert_eq!(messages.len(), 3);
478 assert_eq!(messages[0], r#"{"id":1}"#);
479 assert_eq!(messages[1], r#"{"id":2}"#);
480 assert_eq!(messages[2], r#"{"id":3}"#);
481 }
482
483 #[test]
484 fn test_stdio_transport_shutdown() {
485 let transport = StdioTransport::new();
486 assert!(!transport.is_shutdown());
487
488 transport.shutdown();
489 assert!(transport.is_shutdown());
490 }
491
492 #[test]
493 fn test_stdio_transport_shutdown_handle() {
494 let transport = StdioTransport::new();
495 let handle = transport.shutdown_handle();
496
497 assert!(!handle.load(Ordering::SeqCst));
498 transport.shutdown();
499 assert!(handle.load(Ordering::SeqCst));
500 }
501
502 #[test]
503 fn test_stdio_config() {
504 let config = StdioConfig {
505 max_message_size: 1024,
506 auto_flush: false,
507 stderr_logging: false,
508 buffer_size: 2048,
509 };
510
511 let transport = StdioTransport::with_config(config);
512 assert!(!transport.config.auto_flush);
513 assert!(!transport.config.stderr_logging);
514 }
515
516 #[test]
517 fn test_message_framer_single() {
518 let mut framer = MessageFramer::new();
519 let messages = framer.feed(b"{\"test\":1}\n").unwrap();
520
521 assert_eq!(messages.len(), 1);
522 assert_eq!(messages[0], "{\"test\":1}");
523 assert!(!framer.has_partial());
524 }
525
526 #[test]
527 fn test_message_framer_multiple() {
528 let mut framer = MessageFramer::new();
529 let messages = framer.feed(b"{\"a\":1}\n{\"b\":2}\n").unwrap();
530
531 assert_eq!(messages.len(), 2);
532 assert_eq!(messages[0], "{\"a\":1}");
533 assert_eq!(messages[1], "{\"b\":2}");
534 }
535
536 #[test]
537 fn test_message_framer_partial() {
538 let mut framer = MessageFramer::new();
539
540 let messages1 = framer.feed(b"{\"partial\":").unwrap();
542 assert!(messages1.is_empty());
543 assert!(framer.has_partial());
544
545 let messages2 = framer.feed(b"true}\n").unwrap();
547 assert_eq!(messages2.len(), 1);
548 assert_eq!(messages2[0], "{\"partial\":true}");
549 assert!(!framer.has_partial());
550 }
551
552 #[test]
553 fn test_message_framer_max_size() {
554 let mut framer = MessageFramer::new().with_max_size(10);
555
556 let result = framer.feed(b"this is way too long");
557 assert!(result.is_err());
558 }
559
560 #[test]
561 fn test_frame_unframe() {
562 let original = r#"{"jsonrpc":"2.0"}"#;
563 let framed = frame_message(original);
564 let unframed = unframe_message(&framed);
565
566 assert_eq!(unframed, original);
567 }
568}