Skip to main content

dcp/compat/
stdio.rs

1//! Standard I/O transport for MCP compatibility.
2//!
3//! Provides stdin/stdout handling with line-based JSON-RPC message framing.
4//! Supports async operations, graceful shutdown, and stderr diagnostics.
5
6use 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/// Configuration for stdio transport
17#[derive(Debug, Clone)]
18pub struct StdioConfig {
19    /// Maximum message size (default 10MB)
20    pub max_message_size: usize,
21    /// Whether to flush after each write
22    pub auto_flush: bool,
23    /// Enable stderr logging
24    pub stderr_logging: bool,
25    /// Read buffer size
26    pub buffer_size: usize,
27}
28
29impl Default for StdioConfig {
30    fn default() -> Self {
31        Self {
32            max_message_size: 10 * 1024 * 1024, // 10MB
33            auto_flush: true,
34            stderr_logging: true,
35            buffer_size: 4096,
36        }
37    }
38}
39
40/// Line-based message framing for stdio transport
41pub struct StdioTransport {
42    /// Buffer for reading lines
43    read_buffer: String,
44    /// Configuration
45    config: StdioConfig,
46    /// Shutdown flag
47    shutdown: Arc<AtomicBool>,
48}
49
50impl Default for StdioTransport {
51    fn default() -> Self {
52        Self::new()
53    }
54}
55
56impl StdioTransport {
57    /// Create a new stdio transport
58    pub fn new() -> Self {
59        Self::with_config(StdioConfig::default())
60    }
61
62    /// Create with custom configuration
63    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    /// Set auto-flush behavior
72    pub fn with_auto_flush(mut self, auto_flush: bool) -> Self {
73        self.config.auto_flush = auto_flush;
74        self
75    }
76
77    /// Get shutdown handle for external shutdown signaling
78    pub fn shutdown_handle(&self) -> Arc<AtomicBool> {
79        Arc::clone(&self.shutdown)
80    }
81
82    /// Signal shutdown
83    pub fn shutdown(&self) {
84        self.shutdown.store(true, Ordering::SeqCst);
85    }
86
87    /// Check if shutdown was signaled
88    pub fn is_shutdown(&self) -> bool {
89        self.shutdown.load(Ordering::SeqCst)
90    }
91
92    /// Read a single JSON-RPC message from stdin
93    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                    // EOF - signal graceful shutdown
104                    self.log_stderr("EOF received on stdin, initiating shutdown");
105                    self.shutdown();
106                    return Ok(None);
107                }
108                ReadLineOutcome::Line => {}
109            }
110
111            // Trim trailing newline
112            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    /// Write a JSON-RPC message to stdout
174    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    /// Read messages from stdin until EOF or shutdown
196    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    /// Log a diagnostic message to stderr
208    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    /// Log an error to stderr
215    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    /// Log a debug message to stderr
226    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    /// Format a sanitized diagnostic log line.
237    pub fn format_log_line(prefix: &str, message: &str) -> String {
238        format!("[{}] {}", prefix, sanitize_text(message))
239    }
240}
241
242/// Message framer for handling JSON-RPC over stdio
243pub struct MessageFramer {
244    /// Accumulated buffer for partial messages
245    buffer: Vec<u8>,
246    /// Maximum message size
247    max_size: usize,
248}
249
250impl Default for MessageFramer {
251    fn default() -> Self {
252        Self::new()
253    }
254}
255
256impl MessageFramer {
257    /// Create a new message framer
258    pub fn new() -> Self {
259        Self {
260            buffer: Vec::with_capacity(4096),
261            max_size: 10 * 1024 * 1024, // 10MB
262        }
263    }
264
265    /// Set maximum message size
266    pub fn with_max_size(mut self, max_size: usize) -> Self {
267        self.max_size = max_size;
268        self
269    }
270
271    /// Feed bytes into the framer and extract complete messages
272    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                // Complete message
278                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                // Check size limit
288                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    /// Check if there's a partial message in the buffer
303    pub fn has_partial(&self) -> bool {
304        !self.buffer.is_empty()
305    }
306
307    /// Get the current buffer size
308    pub fn buffer_size(&self) -> usize {
309        self.buffer.len()
310    }
311
312    /// Clear the buffer
313    pub fn clear(&mut self) {
314        self.buffer.clear();
315    }
316}
317
318/// Frame a message for stdio transport (add newline)
319pub fn frame_message(message: &str) -> String {
320    format!("{}\n", message)
321}
322
323/// Unframe a message from stdio transport (remove trailing newline)
324pub fn unframe_message(data: &str) -> &str {
325    data.trim_end_matches('\n').trim_end_matches('\r')
326}
327
328/// Async stdio transport using tokio
329#[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    /// Async stdio transport
336    pub struct AsyncStdioTransport {
337        config: StdioConfig,
338        shutdown: Arc<AtomicBool>,
339    }
340
341    impl AsyncStdioTransport {
342        /// Create new async transport
343        pub fn new(config: StdioConfig) -> Self {
344            Self {
345                config,
346                shutdown: Arc::new(AtomicBool::new(false)),
347            }
348        }
349
350        /// Get shutdown handle
351        pub fn shutdown_handle(&self) -> Arc<AtomicBool> {
352            Arc::clone(&self.shutdown)
353        }
354
355        /// Run the transport, returning channels for messages
356        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            // Spawn stdin reader
364            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                            // EOF
378                            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            // Spawn stdout writer
402            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()); // EOF triggers shutdown
450    }
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        // First chunk - partial message
541        let messages1 = framer.feed(b"{\"partial\":").unwrap();
542        assert!(messages1.is_empty());
543        assert!(framer.has_partial());
544
545        // Second chunk - complete message
546        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}