1use async_trait::async_trait;
7use serde_json::Value;
8use std::collections::HashMap;
9use std::process::Stdio;
10use std::sync::Arc;
11use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, BufWriter};
12use tokio::process::{Child, Command};
13use tokio::sync::{broadcast, mpsc, Mutex, RwLock};
14use tokio::time::{timeout, Duration};
15
16use crate::core::error::{McpError, McpResult};
17use crate::protocol::types::{
18 error_codes, JsonRpcError, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse,
19};
20use crate::protocol::{
21 has_tasks_extension, json_rpc_error_details, methods, modern_request_context,
22 ServerCapabilities, SubscriptionFilter, SubscriptionsAcknowledgedParams,
23 SubscriptionsListenParams, SubscriptionsListenResult, HEADER_MISMATCH,
24 MISSING_REQUIRED_CLIENT_CAPABILITY, SUBSCRIPTION_ID_META_KEY, TASKS_EXTENSION_ID,
25 UNSUPPORTED_PROTOCOL_VERSION,
26};
27use crate::transport::traits::{
28 ClientSubscription, ConnectionState, ServerRequestHandler, ServerTransport, Transport,
29 TransportConfig,
30};
31
32fn add_subscription_id(notification: &mut JsonRpcNotification, id: &Value) {
33 let params = notification
34 .params
35 .get_or_insert_with(|| Value::Object(serde_json::Map::new()));
36 let Some(params) = params.as_object_mut() else {
37 return;
38 };
39 let meta = params
40 .entry("_meta")
41 .or_insert_with(|| Value::Object(serde_json::Map::new()));
42 if let Some(meta) = meta.as_object_mut() {
43 meta.insert(SUBSCRIPTION_ID_META_KEY.to_string(), id.clone());
44 }
45}
46
47async fn write_stdio_line<T: serde::Serialize>(
48 writer: &Arc<Mutex<BufWriter<tokio::io::Stdout>>>,
49 message: &T,
50) -> McpResult<()> {
51 let line = serde_json::to_string(message).map_err(McpError::serialization)?;
52 let mut writer = writer.lock().await;
53 writer
54 .write_all(line.as_bytes())
55 .await
56 .map_err(McpError::io)?;
57 writer.write_all(b"\n").await.map_err(McpError::io)?;
58 writer.flush().await.map_err(McpError::io)
59}
60
61fn accepted_subscription_filter(
62 requested: &SubscriptionFilter,
63 capabilities: &ServerCapabilities,
64) -> SubscriptionFilter {
65 let tasks_enabled = capabilities
66 .extensions
67 .as_ref()
68 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID));
69 SubscriptionFilter {
70 tools_list_changed: (requested.tools_list_changed == Some(true)
71 && capabilities
72 .tools
73 .as_ref()
74 .is_some_and(|value| value.list_changed == Some(true)))
75 .then_some(true),
76 prompts_list_changed: (requested.prompts_list_changed == Some(true)
77 && capabilities
78 .prompts
79 .as_ref()
80 .is_some_and(|value| value.list_changed == Some(true)))
81 .then_some(true),
82 resources_list_changed: (requested.resources_list_changed == Some(true)
83 && capabilities
84 .resources
85 .as_ref()
86 .is_some_and(|value| value.list_changed == Some(true)))
87 .then_some(true),
88 resource_subscriptions: if capabilities
89 .resources
90 .as_ref()
91 .is_some_and(|value| value.subscribe == Some(true))
92 {
93 requested.resource_subscriptions.clone()
94 } else {
95 Vec::new()
96 },
97 task_ids: if tasks_enabled {
98 requested.task_ids.clone()
99 } else {
100 Vec::new()
101 },
102 }
103}
104
105#[derive(Debug)]
110pub struct StdioClientTransport {
111 child: Option<Child>,
112 stdin_writer: Option<BufWriter<tokio::process::ChildStdin>>,
113 #[allow(dead_code)]
114 stdout_reader: Option<BufReader<tokio::process::ChildStdout>>,
115 notification_receiver: Option<mpsc::UnboundedReceiver<JsonRpcNotification>>,
116 pending_requests:
117 Arc<Mutex<HashMap<Value, tokio::sync::oneshot::Sender<McpResult<JsonRpcResponse>>>>>,
118 subscription_senders: Arc<Mutex<HashMap<Value, mpsc::UnboundedSender<JsonRpcNotification>>>>,
119 config: TransportConfig,
120 state: ConnectionState,
121}
122
123impl StdioClientTransport {
124 pub async fn new<S: AsRef<str>>(command: S, args: Vec<S>) -> McpResult<Self> {
133 Self::with_config(command, args, TransportConfig::default()).await
134 }
135
136 pub async fn new_with_command(command: &str, args: &[String]) -> McpResult<Self> {
145 let args_str: Vec<&str> = args.iter().map(|s| s.as_str()).collect();
146 Self::new(command, args_str).await
147 }
148
149 pub async fn with_env<S: AsRef<str>>(
159 command: S,
160 args: Vec<S>,
161 env: HashMap<String, String>,
162 ) -> McpResult<Self> {
163 Self::with_config_and_env(command, args, TransportConfig::default(), Some(env)).await
164 }
165
166 pub async fn with_config<S: AsRef<str>>(
176 command: S,
177 args: Vec<S>,
178 config: TransportConfig,
179 ) -> McpResult<Self> {
180 Self::with_config_and_env(command, args, config, None).await
181 }
182
183 pub async fn with_config_and_env<S: AsRef<str>>(
194 command: S,
195 args: Vec<S>,
196 config: TransportConfig,
197 env: Option<HashMap<String, String>>,
198 ) -> McpResult<Self> {
199 let command_str = command.as_ref();
200 let args_str: Vec<&str> = args.iter().map(|s| s.as_ref()).collect();
201
202 tracing::debug!("Starting MCP server: {} {:?}", command_str, args_str);
203
204 let mut cmd = Command::new(command_str);
205 cmd.args(&args_str)
206 .stdin(Stdio::piped())
207 .stdout(Stdio::piped())
208 .stderr(Stdio::piped());
209
210 if let Some(env_vars) = env {
212 cmd.envs(env_vars);
213 }
214
215 let mut child = cmd
216 .spawn()
217 .map_err(|e| McpError::transport(format!("Failed to start server process: {e}")))?;
218
219 let stdin = child
220 .stdin
221 .take()
222 .ok_or_else(|| McpError::transport("Failed to get stdin handle"))?;
223 let stdout = child
224 .stdout
225 .take()
226 .ok_or_else(|| McpError::transport("Failed to get stdout handle"))?;
227
228 let stdin_writer = BufWriter::new(stdin);
229 let stdout_reader = BufReader::new(stdout);
230
231 let (notification_sender, notification_receiver) = mpsc::unbounded_channel();
232 let pending_requests = Arc::new(Mutex::new(HashMap::new()));
233 let subscription_senders = Arc::new(Mutex::new(HashMap::new()));
234
235 let reader_pending_requests = pending_requests.clone();
237 let reader_subscription_senders = subscription_senders.clone();
238 let reader = stdout_reader;
239 tokio::spawn(async move {
240 Self::message_processor(
241 reader,
242 notification_sender,
243 reader_pending_requests,
244 reader_subscription_senders,
245 )
246 .await;
247 });
248
249 Ok(Self {
250 child: Some(child),
251 stdin_writer: Some(stdin_writer),
252 stdout_reader: None, notification_receiver: Some(notification_receiver),
254 pending_requests,
255 subscription_senders,
256 config,
257 state: ConnectionState::Connected,
258 })
259 }
260
261 async fn message_processor(
262 mut reader: BufReader<tokio::process::ChildStdout>,
263 notification_sender: mpsc::UnboundedSender<JsonRpcNotification>,
264 pending_requests: Arc<
265 Mutex<HashMap<Value, tokio::sync::oneshot::Sender<McpResult<JsonRpcResponse>>>>,
266 >,
267 subscription_senders: Arc<
268 Mutex<HashMap<Value, mpsc::UnboundedSender<JsonRpcNotification>>>,
269 >,
270 ) {
271 let mut line = String::new();
272
273 loop {
274 line.clear();
275 match reader.read_line(&mut line).await {
276 Ok(0) => {
277 tracing::debug!("STDIO reader reached EOF");
278 break;
279 }
280 Ok(_) => {
281 let line = line.trim();
282 if line.is_empty() {
283 continue;
284 }
285
286 tracing::trace!("Received: {}", line);
287
288 let parsed_value = serde_json::from_str::<Value>(line).ok();
289 if parsed_value
290 .as_ref()
291 .and_then(|value| value.get("error"))
292 .is_some()
293 {
294 let Ok(error_response) = serde_json::from_str::<JsonRpcError>(line) else {
295 tracing::warn!("Failed to parse JSON-RPC error: {}", line);
296 continue;
297 };
298 let error = match error_response.error.code {
299 error_codes::METHOD_NOT_FOUND => {
300 McpError::MethodNotFound(error_response.error.message)
301 }
302 HEADER_MISMATCH => {
303 McpError::HeaderMismatch(error_response.error.message)
304 }
305 MISSING_REQUIRED_CLIENT_CAPABILITY => {
306 let required = error_response
307 .error
308 .data
309 .and_then(|value| value.get("requiredCapabilities").cloned())
310 .unwrap_or_else(|| serde_json::json!({}));
311 McpError::MissingRequiredClientCapability(required)
312 }
313 UNSUPPORTED_PROTOCOL_VERSION => {
314 let data = error_response.error.data.unwrap_or_default();
315 McpError::UnsupportedProtocolVersion {
316 requested: data
317 .get("requested")
318 .and_then(Value::as_str)
319 .unwrap_or("unknown")
320 .to_string(),
321 supported: data
322 .get("supported")
323 .and_then(Value::as_array)
324 .into_iter()
325 .flatten()
326 .filter_map(Value::as_str)
327 .map(str::to_string)
328 .collect(),
329 }
330 }
331 code => McpError::Protocol(format!(
332 "JSON-RPC error {code}: {}",
333 error_response.error.message
334 )),
335 };
336 let mut pending = pending_requests.lock().await;
337 if let Some(sender) = pending.remove(&error_response.id) {
338 let _ = sender.send(Err(error));
339 }
340 subscription_senders.lock().await.remove(&error_response.id);
341 }
342 else if let Ok(response) = serde_json::from_str::<JsonRpcResponse>(line) {
344 let mut pending = pending_requests.lock().await;
345 match pending.remove(&response.id) {
346 Some(sender) => {
347 let response_id = response.id.clone();
348 let _ = sender.send(Ok(response));
349 subscription_senders.lock().await.remove(&response_id);
350 }
351 _ => {
352 tracing::warn!(
353 "Received response for unknown request ID: {:?}",
354 response.id
355 );
356 }
357 }
358 }
359 else if let Ok(notification) =
361 serde_json::from_str::<JsonRpcNotification>(line)
362 {
363 if let Some(subscription_id) = notification
364 .params
365 .as_ref()
366 .and_then(|params| params.get("_meta"))
367 .and_then(|meta| meta.get(SUBSCRIPTION_ID_META_KEY))
368 {
369 if let Some(sender) = subscription_senders
370 .lock()
371 .await
372 .get(subscription_id)
373 .cloned()
374 {
375 let _ = sender.send(notification.clone());
376 }
377 }
378 if notification_sender.send(notification).is_err() {
379 tracing::debug!("Notification receiver dropped");
380 break;
381 }
382 } else {
383 tracing::warn!("Failed to parse message: {}", line);
384 }
385 }
386 Err(e) => {
387 tracing::error!("Error reading from stdout: {}", e);
388 break;
389 }
390 }
391 }
392 }
393}
394
395#[async_trait]
396impl Transport for StdioClientTransport {
397 async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
398 let writer = self
399 .stdin_writer
400 .as_mut()
401 .ok_or_else(|| McpError::transport("Transport not connected"))?;
402
403 let (sender, receiver) = tokio::sync::oneshot::channel();
404
405 {
407 let mut pending = self.pending_requests.lock().await;
408 pending.insert(request.id.clone(), sender);
409 }
410
411 let request_line = serde_json::to_string(&request).map_err(McpError::serialization)?;
413
414 tracing::trace!("Sending: {}", request_line);
415
416 writer
417 .write_all(request_line.as_bytes())
418 .await
419 .map_err(|e| McpError::transport(format!("Failed to write request: {e}")))?;
420 writer
421 .write_all(b"\n")
422 .await
423 .map_err(|e| McpError::transport(format!("Failed to write newline: {e}")))?;
424 writer
425 .flush()
426 .await
427 .map_err(|e| McpError::transport(format!("Failed to flush: {e}")))?;
428
429 let timeout_duration = Duration::from_millis(self.config.read_timeout_ms.unwrap_or(60_000));
431
432 let response = timeout(timeout_duration, receiver)
433 .await
434 .map_err(|_| McpError::timeout("Request timeout"))?
435 .map_err(|_| McpError::transport("Response channel closed"))??;
436
437 Ok(response)
438 }
439
440 async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
441 let writer = self
442 .stdin_writer
443 .as_mut()
444 .ok_or_else(|| McpError::transport("Transport not connected"))?;
445
446 let notification_line =
447 serde_json::to_string(¬ification).map_err(McpError::serialization)?;
448
449 tracing::trace!("Sending notification: {}", notification_line);
450
451 writer
452 .write_all(notification_line.as_bytes())
453 .await
454 .map_err(|e| McpError::transport(format!("Failed to write notification: {e}")))?;
455 writer
456 .write_all(b"\n")
457 .await
458 .map_err(|e| McpError::transport(format!("Failed to write newline: {e}")))?;
459 writer
460 .flush()
461 .await
462 .map_err(|e| McpError::transport(format!("Failed to flush: {e}")))?;
463
464 Ok(())
465 }
466
467 async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
468 if let Some(ref mut receiver) = self.notification_receiver {
469 match receiver.try_recv() {
470 Ok(notification) => Ok(Some(notification)),
471 Err(mpsc::error::TryRecvError::Empty) => Ok(None),
472 Err(mpsc::error::TryRecvError::Disconnected) => {
473 Err(McpError::transport("Notification channel disconnected"))
474 }
475 }
476 } else {
477 Ok(None)
478 }
479 }
480
481 async fn open_subscription(
482 &mut self,
483 request: JsonRpcRequest,
484 ) -> McpResult<ClientSubscription> {
485 if request.method != methods::SUBSCRIPTIONS_LISTEN {
486 return Err(McpError::InvalidParams(
487 "open_subscription requires subscriptions/listen".to_string(),
488 ));
489 }
490 let writer = self
491 .stdin_writer
492 .as_mut()
493 .ok_or_else(|| McpError::transport("Transport not connected"))?;
494 let (response_sender, response_receiver) = tokio::sync::oneshot::channel();
495 let (notification_sender, notification_receiver) = mpsc::unbounded_channel();
496 self.pending_requests
497 .lock()
498 .await
499 .insert(request.id.clone(), response_sender);
500 self.subscription_senders
501 .lock()
502 .await
503 .insert(request.id.clone(), notification_sender);
504 let line = serde_json::to_string(&request).map_err(McpError::serialization)?;
505 if let Err(error) = async {
506 writer.write_all(line.as_bytes()).await?;
507 writer.write_all(b"\n").await?;
508 writer.flush().await
509 }
510 .await
511 {
512 self.pending_requests.lock().await.remove(&request.id);
513 self.subscription_senders.lock().await.remove(&request.id);
514 return Err(McpError::io(error));
515 }
516 Ok(ClientSubscription::new(
517 request.id,
518 notification_receiver,
519 response_receiver,
520 ))
521 }
522
523 async fn cancel_subscription(&mut self, request_id: &Value) -> McpResult<()> {
524 let notification = JsonRpcNotification::new(
525 methods::CANCELLED.to_string(),
526 Some(serde_json::json!({"requestId": request_id})),
527 )?;
528 self.send_notification(notification).await
529 }
530
531 async fn close(&mut self) -> McpResult<()> {
532 tracing::debug!("Closing STDIO transport");
533
534 self.state = ConnectionState::Closing;
535
536 if let Some(mut writer) = self.stdin_writer.take() {
538 let _ = writer.shutdown().await;
539 }
540
541 if let Some(mut child) = self.child.take() {
543 match timeout(Duration::from_secs(5), child.wait()).await {
544 Ok(Ok(status)) => {
545 tracing::debug!("Server process exited with status: {}", status);
546 }
547 Ok(Err(e)) => {
548 tracing::warn!("Error waiting for server process: {}", e);
549 }
550 Err(_) => {
551 tracing::warn!("Timeout waiting for server process, killing it");
552 let _ = child.kill().await;
553 }
554 }
555 }
556
557 self.state = ConnectionState::Disconnected;
558 Ok(())
559 }
560
561 fn is_connected(&self) -> bool {
562 matches!(self.state, ConnectionState::Connected)
563 }
564
565 fn connection_info(&self) -> String {
566 let state = &self.state;
567 format!("STDIO transport (state: {state:?})")
568 }
569}
570
571pub struct StdioServerTransport {
576 stdin_reader: Option<BufReader<tokio::io::Stdin>>,
577 stdout_writer: Option<Arc<Mutex<BufWriter<tokio::io::Stdout>>>>,
578 #[allow(dead_code)]
579 config: TransportConfig,
580 running: bool,
581 request_handler: Option<ServerRequestHandler>,
582 capabilities: ServerCapabilities,
583 subscriptions: Arc<RwLock<HashMap<Value, SubscriptionFilter>>>,
584 task_notifications: Option<broadcast::Receiver<JsonRpcNotification>>,
585}
586
587impl StdioServerTransport {
588 pub fn new() -> Self {
593 Self::with_config(TransportConfig::default())
594 }
595
596 pub fn with_config(config: TransportConfig) -> Self {
604 let stdin_reader = BufReader::new(tokio::io::stdin());
605 let stdout_writer = Arc::new(Mutex::new(BufWriter::new(tokio::io::stdout())));
606
607 Self {
608 stdin_reader: Some(stdin_reader),
609 stdout_writer: Some(stdout_writer),
610 config,
611 running: false,
612 request_handler: None,
613 capabilities: ServerCapabilities::default(),
614 subscriptions: Arc::new(RwLock::new(HashMap::new())),
615 task_notifications: None,
616 }
617 }
618}
619
620#[async_trait]
621impl ServerTransport for StdioServerTransport {
622 async fn start(&mut self) -> McpResult<()> {
623 tracing::debug!("Starting STDIO server transport");
624
625 let mut reader = self
626 .stdin_reader
627 .take()
628 .ok_or_else(|| McpError::transport("STDIN reader already taken"))?;
629 let writer = self
630 .stdout_writer
631 .as_ref()
632 .cloned()
633 .ok_or_else(|| McpError::transport("STDOUT writer is unavailable"))?;
634
635 self.running = true;
636 let request_handler = self.request_handler.clone();
637 let subscriptions = self.subscriptions.clone();
638 if let Some(mut task_notifications) = self.task_notifications.take() {
639 let task_writer = writer.clone();
640 let task_subscriptions = subscriptions.clone();
641 tokio::spawn(async move {
642 while let Ok(notification) = task_notifications.recv().await {
643 let entries = task_subscriptions.read().await.clone();
644 for (id, filter) in entries {
645 if !filter.matches(¬ification.method, notification.params.as_ref()) {
646 continue;
647 }
648 let mut notification = notification.clone();
649 add_subscription_id(&mut notification, &id);
650 let _ = write_stdio_line(&task_writer, ¬ification).await;
651 }
652 }
653 });
654 }
655
656 let mut line = String::new();
657 loop {
658 line.clear();
659
660 match reader.read_line(&mut line).await {
661 Ok(0) => {
662 tracing::debug!("STDIN closed, stopping server");
663 break;
664 }
665 Ok(_) => {
666 let line = line.trim();
667 if line.is_empty() {
668 continue;
669 }
670
671 tracing::trace!("Received: {}", line);
672
673 let parsed: Value = match serde_json::from_str(line) {
674 Ok(value) => value,
675 Err(error) => {
676 tracing::warn!(%error, "failed to parse STDIO JSON");
677 continue;
678 }
679 };
680 if parsed.get("method").and_then(Value::as_str) == Some(methods::CANCELLED)
681 && parsed.get("id").is_none()
682 {
683 if let Some(request_id) = parsed
684 .get("params")
685 .and_then(|params| params.get("requestId"))
686 .cloned()
687 {
688 if subscriptions.write().await.remove(&request_id).is_some() {
689 let mut meta = HashMap::new();
690 meta.insert(
691 SUBSCRIPTION_ID_META_KEY.to_string(),
692 request_id.clone(),
693 );
694 let response = JsonRpcResponse::success(
695 request_id,
696 serde_json::to_value(SubscriptionsListenResult {
697 result_type: "complete".to_string(),
698 meta,
699 })?,
700 )?;
701 write_stdio_line(&writer, &response).await?;
702 }
703 }
704 continue;
705 }
706
707 match serde_json::from_value::<JsonRpcRequest>(parsed) {
709 Ok(request) => {
710 if request.method == methods::SUBSCRIPTIONS_LISTEN {
711 let context =
712 modern_request_context(&request)?.ok_or_else(|| {
713 McpError::InvalidParams(
714 "subscriptions/listen requires modern metadata"
715 .to_string(),
716 )
717 })?;
718 let params: SubscriptionsListenParams = serde_json::from_value(
719 request.params.clone().ok_or_else(|| {
720 McpError::InvalidParams(
721 "missing subscription filter".to_string(),
722 )
723 })?,
724 )?;
725 if params.notifications.requests_tasks()
726 && !has_tasks_extension(&context.client_capabilities)
727 {
728 let error = McpError::MissingRequiredClientCapability(
729 serde_json::json!({"extensions": {(TASKS_EXTENSION_ID): {}}}),
730 );
731 let (code, data) = json_rpc_error_details(&error);
732 let response = JsonRpcError::error(
733 request.id,
734 code,
735 error.to_string(),
736 data,
737 );
738 write_stdio_line(&writer, &response).await?;
739 continue;
740 }
741 let accepted = accepted_subscription_filter(
742 ¶ms.notifications,
743 &self.capabilities,
744 );
745 subscriptions
746 .write()
747 .await
748 .insert(request.id.clone(), accepted.clone());
749 let mut meta = HashMap::new();
750 meta.insert(
751 SUBSCRIPTION_ID_META_KEY.to_string(),
752 request.id.clone(),
753 );
754 let acknowledgement = JsonRpcNotification::new(
755 methods::SUBSCRIPTIONS_ACKNOWLEDGED.to_string(),
756 Some(SubscriptionsAcknowledgedParams {
757 notifications: accepted,
758 meta,
759 }),
760 )?;
761 write_stdio_line(&writer, &acknowledgement).await?;
762 continue;
763 }
764 let response_result = if let Some(ref handler) = request_handler {
765 handler(request.clone()).await
767 } else {
768 Err(McpError::protocol(format!(
770 "Method '{}' not found",
771 request.method
772 )))
773 };
774
775 let response_or_error = match response_result {
776 Ok(response) => serde_json::to_string(&response),
777 Err(error) => {
778 let (code, data) = json_rpc_error_details(&error);
780 let json_rpc_error = crate::protocol::types::JsonRpcError {
781 jsonrpc: "2.0".to_string(),
782 id: request.id,
783 error: crate::protocol::types::ErrorObject {
784 code,
785 message: error.to_string(),
786 data,
787 },
788 };
789 serde_json::to_string(&json_rpc_error)
790 }
791 };
792
793 let response_line =
794 response_or_error.map_err(McpError::serialization)?;
795
796 tracing::trace!("Sending: {}", response_line);
797
798 let mut writer_guard = writer.lock().await;
799 writer_guard
800 .write_all(response_line.as_bytes())
801 .await
802 .map_err(|e| {
803 McpError::transport(format!("Failed to write response: {e}"))
804 })?;
805 writer_guard.write_all(b"\n").await.map_err(|e| {
806 McpError::transport(format!("Failed to write newline: {e}"))
807 })?;
808 writer_guard.flush().await.map_err(|e| {
809 McpError::transport(format!("Failed to flush: {e}"))
810 })?;
811 }
812 Err(e) => {
813 tracing::warn!("Failed to parse request: {} - Error: {}", line, e);
814 }
817 }
818 }
819 Err(e) => {
820 tracing::error!("Error reading from stdin: {}", e);
821 return Err(McpError::io(e));
822 }
823 }
824 }
825
826 Ok(())
827 }
828
829 fn set_request_handler(&mut self, handler: ServerRequestHandler) {
830 self.request_handler = Some(handler);
831 }
832
833 fn set_server_capabilities(&mut self, capabilities: ServerCapabilities) -> McpResult<()> {
834 self.capabilities = capabilities;
835 Ok(())
836 }
837
838 fn set_task_notifications(
839 &mut self,
840 receiver: broadcast::Receiver<JsonRpcNotification>,
841 ) -> McpResult<()> {
842 self.task_notifications = Some(receiver);
843 Ok(())
844 }
845
846 async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
847 let writer = self
848 .stdout_writer
849 .as_ref()
850 .ok_or_else(|| McpError::transport("STDOUT writer not available"))?;
851
852 let subscriptions = self.subscriptions.read().await.clone();
853 if subscriptions.is_empty() {
854 tracing::trace!(method = %notification.method, "sending legacy notification");
856 write_stdio_line(writer, ¬ification).await?;
857 } else {
858 for (id, filter) in subscriptions {
859 if !filter.matches(¬ification.method, notification.params.as_ref()) {
860 continue;
861 }
862 let mut routed = notification.clone();
863 add_subscription_id(&mut routed, &id);
864 write_stdio_line(writer, &routed).await?;
865 }
866 }
867
868 Ok(())
869 }
870
871 async fn stop(&mut self) -> McpResult<()> {
872 tracing::debug!("Stopping STDIO server transport");
873 self.running = false;
874 Ok(())
875 }
876
877 fn is_running(&self) -> bool {
878 self.running
879 }
880
881 fn server_info(&self) -> String {
882 format!("STDIO server transport (running: {})", self.running)
883 }
884}
885
886impl StdioServerTransport {
888 pub async fn handle_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
891 Err(McpError::protocol(format!(
893 "Method '{}' not found (test mode)",
894 request.method
895 )))
896 }
897}
898
899impl Default for StdioServerTransport {
900 fn default() -> Self {
901 Self::new()
902 }
903}
904
905impl Drop for StdioClientTransport {
906 fn drop(&mut self) {
907 if let Some(mut child) = self.child.take() {
908 let _ = child.start_kill();
910 }
911 }
912}
913
914#[cfg(test)]
915mod tests {
916 use super::*;
917 use serde_json::json;
918 use std::collections::HashMap;
919 use std::sync::Arc;
920 use tokio::sync::{mpsc, Mutex};
921
922 #[test]
923 fn test_stdio_server_creation() {
924 let transport = StdioServerTransport::new();
925 assert!(!transport.is_running());
926 assert!(transport.stdin_reader.is_some());
927 assert!(transport.stdout_writer.is_some());
928 }
929
930 #[test]
931 fn test_stdio_server_with_config() {
932 let config = TransportConfig {
933 read_timeout_ms: Some(30_000),
934 ..Default::default()
935 };
936
937 let transport = StdioServerTransport::with_config(config);
938 assert_eq!(transport.config.read_timeout_ms, Some(30_000));
939 }
940
941 #[tokio::test]
942 async fn test_stdio_server_handle_request() {
943 let mut transport = StdioServerTransport::new();
944
945 let request = JsonRpcRequest {
946 jsonrpc: "2.0".to_string(),
947 id: json!(1),
948 method: "unknown_method".to_string(),
949 params: None,
950 };
951
952 let result = transport.handle_request(request).await;
953 assert!(result.is_err());
954
955 match result.unwrap_err() {
956 McpError::Protocol(msg) => assert!(msg.contains("unknown_method")),
957 _ => panic!("Expected Protocol error"),
958 }
959 }
960
961 #[tokio::test]
966 async fn test_client_transport_creation_failure() {
967 let result = StdioClientTransport::new("/nonexistent/command", vec!["arg1"]).await;
969 assert!(result.is_err());
970 match result.unwrap_err() {
971 McpError::Transport(msg) => assert!(msg.contains("Failed to start server process")),
972 _ => panic!("Expected Transport error"),
973 }
974 }
975
976 #[tokio::test]
977 async fn test_client_transport_with_config() {
978 let config = TransportConfig {
979 read_timeout_ms: Some(5000),
980 max_message_size: Some(2048),
981 ..Default::default()
982 };
983
984 let result = StdioClientTransport::with_config("echo", vec!["test"], config.clone()).await;
986
987 if let Ok(transport) = result {
990 assert_eq!(transport.config.read_timeout_ms, Some(5000));
991 assert_eq!(transport.config.max_message_size, Some(2048));
992 }
993 }
994
995 #[tokio::test]
996 async fn test_client_send_request_disconnected() {
997 let mut transport = StdioClientTransport {
998 child: None,
999 stdin_writer: None,
1000 stdout_reader: None,
1001 notification_receiver: None,
1002 pending_requests: Arc::new(Mutex::new(HashMap::new())),
1003 subscription_senders: Arc::new(Mutex::new(HashMap::new())),
1004 config: TransportConfig::default(),
1005 state: ConnectionState::Disconnected,
1006 };
1007
1008 let request = JsonRpcRequest {
1009 jsonrpc: "2.0".to_string(),
1010 id: json!(1),
1011 method: "test_method".to_string(),
1012 params: None,
1013 };
1014
1015 let result = transport.send_request(request).await;
1016 assert!(result.is_err());
1017 match result.unwrap_err() {
1018 McpError::Transport(msg) => assert!(msg.contains("not connected")),
1019 _ => panic!("Expected Transport error"),
1020 }
1021 }
1022
1023 #[tokio::test]
1024 async fn test_client_receive_notification() {
1025 let (tx, rx) = mpsc::unbounded_channel();
1026
1027 let mut transport = StdioClientTransport {
1028 child: None,
1029 stdin_writer: None,
1030 stdout_reader: None,
1031 notification_receiver: Some(rx),
1032 pending_requests: Arc::new(Mutex::new(HashMap::new())),
1033 subscription_senders: Arc::new(Mutex::new(HashMap::new())),
1034 config: TransportConfig::default(),
1035 state: ConnectionState::Connected,
1036 };
1037
1038 let notification = JsonRpcNotification {
1040 jsonrpc: "2.0".to_string(),
1041 method: "test_notification".to_string(),
1042 params: Some(json!({"test": true})),
1043 };
1044 tx.send(notification.clone()).unwrap();
1045
1046 let received = transport.receive_notification().await.unwrap();
1047 assert_eq!(received.unwrap().method, "test_notification");
1048 }
1049
1050 #[tokio::test]
1051 async fn test_client_receive_notification_timeout() {
1052 let (_tx, rx) = mpsc::unbounded_channel();
1053
1054 let mut transport = StdioClientTransport {
1055 child: None,
1056 stdin_writer: None,
1057 stdout_reader: None,
1058 notification_receiver: Some(rx),
1059 pending_requests: Arc::new(Mutex::new(HashMap::new())),
1060 subscription_senders: Arc::new(Mutex::new(HashMap::new())),
1061 config: TransportConfig {
1062 read_timeout_ms: Some(100),
1063 ..Default::default()
1064 },
1065 state: ConnectionState::Connected,
1066 };
1067
1068 let result = transport.receive_notification().await;
1069 assert!(result.is_ok());
1071 assert!(result.unwrap().is_none());
1072 }
1073
1074 }