1pub mod mcp_server_runtime;
2pub mod mcp_server_runtime_core;
3use crate::auth::AuthInfo;
4use crate::error::SdkResult;
5use crate::mcp_traits::{
6 McpObserver, McpServer, McpServerHandler, RequestIdGen, RequestIdGenNumeric,
7};
8use crate::schema::{
9 schema_utils::{
10 ClientMessage, ClientMessages, FromMessage, MessageFromServer, SdkError, ServerMessage,
11 ServerMessages,
12 },
13 InitializeRequestParams, InitializeResult, RequestId, RpcError,
14};
15use crate::task_store::{ClientTaskStore, ServerTaskStore, TaskStatusPoller, TaskStatusUpdate};
16use crate::utils::AbortTaskOnDrop;
17use async_trait::async_trait;
18use futures::future::try_join_all;
19use futures::{StreamExt, TryFutureExt};
20use rust_mcp_schema::{GetTaskParams, GetTaskPayloadParams};
21use rust_mcp_transport::SessionId;
22use rust_mcp_transport::{IoStream, TaskId, TransportDispatcher};
23use std::panic;
24use std::sync::Arc;
25use std::time::Duration;
26use tokio::io::AsyncWriteExt;
27use tokio::sync::{mpsc, oneshot, watch, RwLock, RwLockReadGuard};
28
29pub const DEFAULT_STREAM_ID: &str = "STANDALONE-STREAM";
30const TASK_CHANNEL_CAPACITY: usize = 500;
31
32tokio::task_local! {
33 pub(crate) static ACTIVE_REQUEST_TRANSPORT: TransportType;
37}
38
39type TransportType = Arc<
41 dyn TransportDispatcher<
42 ClientMessages,
43 MessageFromServer,
44 ClientMessage,
45 ServerMessages,
46 ServerMessage,
47 >,
48>;
49
50pub struct ServerRuntime {
52 handler: Arc<dyn McpServerHandler>,
54 server_details: Arc<InitializeResult>,
56 session_id: Option<SessionId>,
57 transport_map: tokio::sync::RwLock<Option<TransportType>>,
58 request_id_gen: Box<dyn RequestIdGen>,
59 client_details_tx: watch::Sender<Option<InitializeRequestParams>>,
60 client_details_rx: watch::Receiver<Option<InitializeRequestParams>>,
61 auth_info: tokio::sync::RwLock<Option<AuthInfo>>,
62 task_store: Option<Arc<ServerTaskStore>>,
63 client_task_store: Option<Arc<ClientTaskStore>>,
64 message_observer: Option<Arc<dyn McpObserver<ClientMessage, ServerMessage>>>,
65}
66
67pub struct McpServerOptions<T>
68where
69 T: TransportDispatcher<
70 ClientMessages,
71 MessageFromServer,
72 ClientMessage,
73 ServerMessages,
74 ServerMessage,
75 >,
76{
77 pub server_details: InitializeResult,
78 pub transport: T,
79 pub handler: Arc<dyn McpServerHandler>,
80 pub task_store: Option<Arc<ServerTaskStore>>,
81 pub client_task_store: Option<Arc<ClientTaskStore>>,
82 pub message_observer: Option<Arc<dyn McpObserver<ClientMessage, ServerMessage>>>,
83}
84
85#[async_trait]
86impl McpServer for ServerRuntime {
87 fn task_store(&self) -> Option<Arc<ServerTaskStore>> {
88 self.task_store.clone()
89 }
90
91 fn client_task_store(&self) -> Option<Arc<ClientTaskStore>> {
92 self.client_task_store.clone()
93 }
94
95 async fn set_client_details(&self, client_details: InitializeRequestParams) -> SdkResult<()> {
97 self.client_details_tx
98 .send(Some(client_details))
99 .map_err(|_| {
100 RpcError::internal_error()
101 .with_message("Failed to set client details".to_string())
102 .into()
103 })
104 }
105
106 async fn update_auth_info(&self, new_auth_info: Option<AuthInfo>) {
107 let should_update = {
108 let current = self.auth_info.read().await;
109 match (&*current, &new_auth_info) {
110 (None, Some(_)) => true,
111 (Some(old), Some(new)) => old.token_unique_id != new.token_unique_id,
112 (Some(_), None) => true,
113 (None, None) => false,
114 }
115 };
116
117 if should_update {
118 *self.auth_info.write().await = new_auth_info;
119 }
120 }
121
122 async fn auth_info(&self) -> RwLockReadGuard<'_, Option<AuthInfo>> {
123 self.auth_info.read().await
124 }
125 async fn auth_info_cloned(&self) -> Option<AuthInfo> {
126 let guard = self.auth_info.read().await;
127 guard.clone()
128 }
129
130 async fn wait_for_initialization(&self) {
131 loop {
132 if self.client_details_rx.borrow().is_some() {
133 return;
134 }
135 let mut rx = self.client_details_rx.clone();
136 rx.changed().await.ok();
137 }
138 }
139
140 async fn send(
141 &self,
142 message: MessageFromServer,
143 request_id: Option<RequestId>,
144 request_timeout: Option<Duration>,
145 ) -> SdkResult<Option<ClientMessage>> {
146 let outgoing_request_id = self
147 .request_id_gen
148 .request_id_for_message(&message, request_id);
149
150 let is_notification = matches!(&message, MessageFromServer::NotificationFromServer(_));
155
156 if is_notification {
157 if let Ok(req_transport) = ACTIVE_REQUEST_TRANSPORT.try_with(|t| t.clone()) {
158 let mcp_message = ServerMessage::from_message(message, outgoing_request_id)?;
159 if let Some(observer) = self.message_observer.as_ref() {
160 observer.on_send(&mcp_message);
161 }
162 return Ok(req_transport
163 .send_message(ServerMessages::Single(mcp_message), request_timeout)
164 .await?
165 .map(|res| res.as_single())
166 .transpose()?);
167 }
168 }
169
170 let mcp_message = ServerMessage::from_message(message, outgoing_request_id)?;
171 if let Some(observer) = self.message_observer.as_ref() {
172 observer.on_send(&mcp_message);
173 }
174
175 let transport_map = self.transport_map.read().await;
176 let transport = transport_map.as_ref().ok_or(
177 RpcError::internal_error()
178 .with_message("transport stream does not exists or is closed!".to_string()),
179 )?;
180
181 let response = transport
182 .send_message(ServerMessages::Single(mcp_message), request_timeout)
183 .await?
184 .map(|res| res.as_single())
185 .transpose()?;
186
187 Ok(response)
188 }
189
190 async fn send_batch(
191 &self,
192 messages: Vec<ServerMessage>,
193 request_timeout: Option<Duration>,
194 ) -> SdkResult<Option<Vec<ClientMessage>>> {
195 let transport_map = self.transport_map.read().await;
196 let transport = transport_map.as_ref().ok_or(
197 RpcError::internal_error()
198 .with_message("transport stream does not exists or is closed!".to_string()),
199 )?;
200
201 if let Some(observer) = self.message_observer.as_ref() {
203 messages.iter().for_each(|msg| observer.on_send(msg));
204 }
205
206 transport
207 .send_batch(messages, request_timeout)
208 .map_err(|err| err.into())
209 .await
210 }
211
212 fn server_info(&self) -> &InitializeResult {
215 &self.server_details
216 }
217
218 fn client_info(&self) -> Option<InitializeRequestParams> {
220 self.client_details_rx.borrow().clone()
221 }
222
223 async fn start(self: Arc<Self>) -> SdkResult<()> {
225 let self_clone = self.clone();
226 let transport_map = self_clone.transport_map.read().await;
227
228 let transport = transport_map.as_ref().ok_or(
229 RpcError::internal_error()
230 .with_message("transport stream does not exists or is closed!".to_string()),
231 )?;
232
233 let mut stream = transport.start().await?;
234
235 let (tx, mut rx) = mpsc::channel(TASK_CHANNEL_CAPACITY);
237
238 while let Some(mcp_messages) = stream.next().await {
240 match mcp_messages {
241 ClientMessages::Single(client_message) => {
242 let transport = transport.clone();
243 let self = self.clone();
244 let tx = tx.clone();
245
246 tokio::spawn(async move {
248 let result = self.handle_message(client_message, &transport).await;
249
250 let send_result: SdkResult<_> = match result {
251 Ok(result) => {
252 if let Some(result) = result {
253 transport
254 .send_message(ServerMessages::Single(result), None)
255 .map_err(|e| e.into())
256 .await
257 } else {
258 Ok(None)
259 }
260 }
261 Err(error) => {
262 tracing::error!("Error handling message : {}", error);
263 Ok(None)
264 }
265 };
266 if let Err(error) = tx.send(send_result).await {
268 tracing::error!("Failed to send result to channel: {}", error);
269 }
270 });
271 }
272 ClientMessages::Batch(client_messages) => {
273 let transport = transport.clone();
274 let self = self_clone.clone();
275 let tx = tx.clone();
276
277 tokio::spawn(async move {
278 let handling_tasks: Vec<_> = client_messages
279 .into_iter()
280 .map(|client_message| self.handle_message(client_message, &transport))
281 .collect();
282
283 let send_result = match try_join_all(handling_tasks).await {
284 Ok(results) => {
285 let results: Vec<_> = results.into_iter().flatten().collect();
286 if !results.is_empty() {
287 transport
288 .send_message(ServerMessages::Batch(results), None)
289 .map_err(|e| e.into())
290 .await
291 } else {
292 Ok(None)
293 }
294 }
295 Err(error) => Err(error),
296 };
297
298 if let Err(error) = tx.send(send_result).await {
299 tracing::error!("Failed to send batch result to channel: {}", error);
300 }
301 });
302 }
303 }
304
305 while let Ok(result) = rx.try_recv() {
307 result?; }
309 }
310
311 drop(tx);
313 while let Some(result) = rx.recv().await {
314 result?; }
316
317 return Ok(());
318 }
319
320 async fn stderr_message(&self, message: String) -> SdkResult<()> {
321 let transport_map = self.transport_map.read().await;
322 let transport = transport_map.as_ref().ok_or(
323 RpcError::internal_error()
324 .with_message("transport stream does not exists or is closed!".to_string()),
325 )?;
326 let mut lock = transport.error_stream().write().await;
327
328 if let Some(IoStream::Writable(stderr)) = lock.as_mut() {
329 stderr.write_all(message.as_bytes()).await?;
330 stderr.write_all(b"\n").await?;
331 stderr.flush().await?;
332 }
333 Ok(())
334 }
335
336 fn session_id(&self) -> Option<SessionId> {
337 self.session_id.to_owned()
338 }
339}
340
341impl ServerRuntime {
342 pub(crate) async fn consume_payload_string(&self, payload: &str) -> SdkResult<()> {
343 let transport_map = self.transport_map.read().await;
344
345 let transport = transport_map.as_ref().ok_or(
346 RpcError::internal_error()
347 .with_message("stream id does not exists or is closed!".to_string()),
348 )?;
349
350 transport.consume_string_payload(payload).await?;
351
352 Ok(())
353 }
354
355 pub(crate) async fn handle_message(
356 self: &Arc<Self>,
357 message: ClientMessage,
358 transport: &Arc<
359 dyn TransportDispatcher<
360 ClientMessages,
361 MessageFromServer,
362 ClientMessage,
363 ServerMessages,
364 ServerMessage,
365 >,
366 >,
367 ) -> SdkResult<Option<ServerMessage>> {
368 if let Some(observer) = self.message_observer.as_ref() {
370 observer.on_receive(&message);
371 }
372
373 let response = match message {
374 ClientMessage::Request(client_jsonrpc_request) => {
376 let request_id = client_jsonrpc_request.request_id().clone();
377
378 let result = self
379 .handler
380 .handle_request(client_jsonrpc_request, self.clone())
381 .await;
382
383 let response: MessageFromServer = match result {
385 Ok(success_value) => success_value.into(),
386 Err(error_value) => {
387 if !self.is_initialized() {
390 return Err(error_value.into());
391 }
392 MessageFromServer::Error(error_value)
393 }
394 };
395
396 let mpc_message: ServerMessage =
397 ServerMessage::from_message(response, Some(request_id))?;
398
399 Some(mpc_message)
400 }
401 ClientMessage::Notification(client_jsonrpc_notification) => {
402 self.handler
403 .handle_notification(client_jsonrpc_notification, self.clone())
404 .await?;
405 None
406 }
407 ClientMessage::Error(jsonrpc_error) => {
408 self.handler
409 .handle_error(&jsonrpc_error.error, self.clone())
410 .await?;
411
412 if let Some(request_id) = jsonrpc_error.id.as_ref() {
413 if let Some(tx_response) = transport.pending_request_tx(request_id).await {
414 tx_response
415 .send(ClientMessage::Error(jsonrpc_error))
416 .map_err(|e| RpcError::internal_error().with_message(e.to_string()))?;
417 } else {
418 tracing::warn!(
419 "Received an error response with no corresponding request {:?}",
420 &jsonrpc_error.id
421 );
422 }
423 }
424 None
425 }
426 ClientMessage::Response(response) => {
427 if let Some(tx_response) = transport.pending_request_tx(&response.id).await {
428 tx_response
429 .send(ClientMessage::Response(response))
430 .map_err(|e| RpcError::internal_error().with_message(e.to_string()))?;
431 } else {
432 tracing::warn!(
433 "Received a response with no corresponding request: {:?}",
434 &response.id
435 );
436 }
437 None
438 }
439 };
440 Ok(response)
441 }
442
443 pub(crate) async fn store_transport(
444 &self,
445 stream_id: &str,
446 transport: Arc<
447 dyn TransportDispatcher<
448 ClientMessages,
449 MessageFromServer,
450 ClientMessage,
451 ServerMessages,
452 ServerMessage,
453 >,
454 >,
455 ) -> SdkResult<()> {
456 if stream_id != DEFAULT_STREAM_ID {
457 return Ok(());
458 }
459 let mut transport_map = self.transport_map.write().await;
460 tracing::trace!("save transport for stream id : {}", stream_id);
461 *transport_map = Some(transport);
462 Ok(())
463 }
464
465 pub(crate) async fn remove_transport(&self, stream_id: &str) -> SdkResult<()> {
467 if stream_id != DEFAULT_STREAM_ID {
468 return Ok(());
469 }
470 let transport_map = self.transport_map.read().await;
471 tracing::trace!("removing transport for stream id : {}", stream_id);
472 if let Some(transport) = transport_map.as_ref() {
473 transport.shut_down().await?;
474 }
475 Ok(())
477 }
478
479 pub(crate) async fn shutdown(&self) {
480 let mut transport_map = self.transport_map.write().await;
481 let transport_option = transport_map.take();
482 drop(transport_map);
483 if let Some(transport) = transport_option {
484 let _ = transport.shut_down().await;
485 }
486 }
487
488 pub(crate) async fn default_stream_exists(&self) -> bool {
489 let transport_map = self.transport_map.read().await;
490 let live_transport = if let Some(t) = transport_map.as_ref() {
491 !t.is_shut_down().await
492 } else {
493 false
494 };
495 live_transport
496 }
497
498 pub(crate) async fn start_stream(
499 self: Arc<Self>,
500 transport: Arc<
501 dyn TransportDispatcher<
502 ClientMessages,
503 MessageFromServer,
504 ClientMessage,
505 ServerMessages,
506 ServerMessage,
507 >,
508 >,
509 stream_id: &str,
510 ping_interval: Duration,
511 payload: Option<String>,
512 ) -> SdkResult<()> {
513 let mut stream = transport.start().await?;
514
515 if stream_id == DEFAULT_STREAM_ID {
516 self.store_transport(stream_id, transport.clone()).await?;
517 }
518
519 let self_clone = self.clone();
520
521 let (disconnect_tx, mut disconnect_rx) = oneshot::channel::<()>();
522 let abort_alive_task = transport
523 .keep_alive(ping_interval, disconnect_tx)
524 .await?
525 .abort_handle();
526
527 let _abort_guard = AbortTaskOnDrop {
529 handle: abort_alive_task,
530 };
531
532 if let Some(payload) = payload {
535 if let Err(err) = transport.consume_string_payload(&payload).await {
536 let _ = self.remove_transport(stream_id).await;
537 return Err(err.into());
538 }
539 }
540
541 let (tx, mut rx) = mpsc::channel(TASK_CHANNEL_CAPACITY);
543
544 loop {
545 tokio::select! {
546 Some(mcp_messages) = stream.next() =>{
547
548 match mcp_messages {
549 ClientMessages::Single(client_message) => {
550 let transport = transport.clone();
551 let self_clone = self.clone();
552 let tx = tx.clone();
553 tokio::spawn(ACTIVE_REQUEST_TRANSPORT.scope(transport.clone(), async move {
554
555 let result = self_clone.handle_message(client_message, &transport).await;
556
557 let send_result: SdkResult<_> = match result {
558 Ok(result) => {
559 if let Some(result) = result {
560 transport
561 .send_message(ServerMessages::Single(result), None)
562 .map_err(|e| e.into())
563 .await
564 } else {
565 Ok(None)
566 }
567 }
568 Err(error) => {
569 tracing::error!("Error handling message : {}", error);
570 Ok(None)
571 }
572 };
573 if let Err(error) = tx.send(send_result).await {
574 tracing::error!("Failed to send batch result to channel: {}", error);
575 }
576 }));
577 }
578 ClientMessages::Batch(client_messages) => {
579
580 let transport = transport.clone();
581 let self_clone = self_clone.clone();
582 let tx = tx.clone();
583
584 tokio::spawn(ACTIVE_REQUEST_TRANSPORT.scope(transport.clone(), async move {
585 let handling_tasks: Vec<_> = client_messages
586 .into_iter()
587 .map(|client_message| self_clone.handle_message(client_message, &transport))
588 .collect();
589
590 let send_result = match try_join_all(handling_tasks).await {
591 Ok(results) => {
592 let results: Vec<_> = results.into_iter().flatten().collect();
593 if !results.is_empty() {
594 transport.send_message(ServerMessages::Batch(results), None)
595 .map_err(|e| e.into())
596 .await
597 }else {
598 Ok(None)
599 }
600 },
601 Err(error) => Err(error),
602 };
603 if let Err(error) = tx.send(send_result).await {
604 tracing::error!("Failed to send batch result to channel: {}", error);
605 }
606 }));
607 }
608 }
609
610 while let Ok(result) = rx.try_recv() {
612 result?; }
614
615 if !stream_id.eq(DEFAULT_STREAM_ID){
617 drop(tx);
618 while let Some(result) = rx.recv().await {
619 result?; }
621 return Ok(());
622 }
623 }
624 _ = &mut disconnect_rx => {
625 drop(tx);
627 while let Some(result) = rx.recv().await {
628 result?; }
630 self.remove_transport(stream_id).await?;
631 return Err(SdkError::connection_closed().into());
633
634 }
635 }
636 }
637 }
638
639 pub(crate) fn new_instance(
640 server_details: Arc<InitializeResult>,
641 handler: Arc<dyn McpServerHandler>,
642 session_id: SessionId,
643 auth_info: Option<AuthInfo>,
644 task_store: Option<Arc<ServerTaskStore>>,
645 client_task_store: Option<Arc<ClientTaskStore>>,
646 message_observer: Option<Arc<dyn McpObserver<ClientMessage, ServerMessage>>>,
647 ) -> Arc<Self> {
648 use tokio::sync::RwLock;
649
650 let (client_details_tx, client_details_rx) =
651 watch::channel::<Option<InitializeRequestParams>>(None);
652 Arc::new(Self {
653 server_details,
654 handler,
655 session_id: Some(session_id),
656 transport_map: tokio::sync::RwLock::new(None),
657 client_details_tx,
658 client_details_rx,
659 request_id_gen: Box::new(RequestIdGenNumeric::new(None)),
660 auth_info: RwLock::new(auth_info),
661 task_store,
662 client_task_store,
663 message_observer,
664 })
665 }
666
667 pub async fn poll_task_status(
668 self: Arc<ServerRuntime>,
669 task_id: TaskId,
670 session_id: Option<String>,
671 task_store: Arc<ClientTaskStore>,
672 ) -> SdkResult<TaskStatusUpdate> {
673 let result = self
674 .request_get_task(GetTaskParams {
675 task_id: task_id.to_string(),
676 })
677 .await?;
678
679 if result.is_terminal() {
680 let task_payload = self
681 .request_get_task_payload(GetTaskPayloadParams {
682 task_id: task_id.clone(),
683 })
684 .await?;
685
686 task_store
687 .store_task_result(
688 task_id.as_str(),
689 result.status,
690 task_payload.into(),
691 session_id.as_ref(),
692 )
693 .await;
694 }
695 Ok((result.status, result.poll_interval))
696 }
697
698 pub(crate) fn new<T>(options: McpServerOptions<T>) -> Arc<Self>
699 where
700 T: TransportDispatcher<
701 ClientMessages,
702 MessageFromServer,
703 ClientMessage,
704 ServerMessages,
705 ServerMessage,
706 >,
707 {
708 let (client_details_tx, client_details_rx) =
709 watch::channel::<Option<InitializeRequestParams>>(None);
710
711 let runtime = Arc::new(Self {
712 server_details: Arc::new(options.server_details),
713 handler: options.handler,
714 session_id: None,
715 transport_map: tokio::sync::RwLock::new(Some(Arc::new(options.transport))),
716 client_details_tx,
717 client_details_rx,
718 request_id_gen: Box::new(RequestIdGenNumeric::new(None)),
719 auth_info: RwLock::new(None),
720 task_store: options.task_store,
721 client_task_store: options.client_task_store,
722 message_observer: options.message_observer,
723 });
724
725 let runtime_clone = runtime.clone();
726 if let Some(task_store) = runtime_clone.task_store() {
727 if let Some(mut stream) = task_store.subscribe() {
729 tokio::spawn(async move {
730 while let Some((params, _)) = stream.next().await {
731 let _ = runtime_clone.notify_task_status(params).await;
732 }
733 });
734 }
735 }
736
737 if let Some(client_task_store) = runtime.client_task_store.clone() {
739 let task_store_clone = client_task_store.clone();
740 let runtime_clone = runtime.clone();
741
742 let callback: TaskStatusPoller = Box::new(move |task_id, session_id| {
743 let task_store_clone = client_task_store.clone();
744 let runtime_clone = runtime_clone.clone();
745
746 Box::pin(async move {
747 runtime_clone
748 .poll_task_status(task_id, session_id, task_store_clone)
749 .await
750 })
751 });
752
753 if let Err(error) = task_store_clone.start_task_polling(callback) {
754 tracing::error!("Failed to start task polling: {error}");
755 }
756 }
757
758 runtime
759 }
760}