Skip to main content

rust_mcp_sdk/mcp_runtimes/
server_runtime.rs

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    /// Per-request transport for sending notifications on the POST response SSE stream.
34    /// Set via `scope()` in spawned handler tasks. Read by `send()` for notification routing.
35    /// Falls back to the GET standalone stream when not set (background tasks, on_initialized, etc.).
36    pub(crate) static ACTIVE_REQUEST_TRANSPORT: TransportType;
37}
38
39// Define a type alias for the TransportDispatcher trait object
40type TransportType = Arc<
41    dyn TransportDispatcher<
42        ClientMessages,
43        MessageFromServer,
44        ClientMessage,
45        ServerMessages,
46        ServerMessage,
47    >,
48>;
49
50/// Struct representing the runtime core of the MCP server, handling transport and client details
51pub struct ServerRuntime {
52    // The handler for processing MCP messages
53    handler: Arc<dyn McpServerHandler>,
54    // Information about the server
55    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    /// Set the client details, storing them in client_details
96    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        // For notifications during a request (tool call), route through the
151        // active POST response stream so the client receives them during
152        // `request()`. Fall back to the GET standalone stream if there is no
153        // active POST stream.
154        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        // telemetry
202        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    /// Returns the server's details, including server capability,
213    /// instructions, protocol_version , server_info and optional meta data
214    fn server_info(&self) -> &InitializeResult {
215        &self.server_details
216    }
217
218    /// Returns the client information if available, after successful initialization , otherwise returns None
219    fn client_info(&self) -> Option<InitializeRequestParams> {
220        self.client_details_rx.borrow().clone()
221    }
222
223    /// Main runtime loop, processes incoming messages and handles requests
224    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        // Create a channel to collect results from spawned tasks
236        let (tx, mut rx) = mpsc::channel(TASK_CHANNEL_CAPACITY);
237
238        // Process incoming messages from the client
239        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                    // Handle incoming messages in a separate task to avoid blocking the stream.
247                    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                        // Send result to the main loop
267                        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            // Check for results from spawned tasks to propagate errors
306            while let Ok(result) = rx.try_recv() {
307                result?; // Propagate errors
308            }
309        }
310
311        // Drop tx to close the channel and collect remaining results
312        drop(tx);
313        while let Some(result) = rx.recv().await {
314            result?; // Propagate errors
315        }
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        // telemetry
369        if let Some(observer) = self.message_observer.as_ref() {
370            observer.on_receive(&message);
371        }
372
373        let response = match message {
374            // Handle a client request
375            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                // create a response to send back to the client
384                let response: MessageFromServer = match result {
385                    Ok(success_value) => success_value.into(),
386                    Err(error_value) => {
387                        // Error occurred during initialization.
388                        // A likely cause could be an unsupported protocol version.
389                        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    //TODO: re-visit and simplify unnecessary hashmap
466    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        // transport_map.remove(stream_id);
476        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        // ensure keep_alive task will be aborted
528        let _abort_guard = AbortTaskOnDrop {
529            handle: abort_alive_task,
530        };
531
532        // in case there is a payload, we consume it by transport to get processed
533        // payload would be message payload coming from the client
534        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        // Create a channel to collect results from spawned tasks
542        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                    // Check for results from spawned tasks to propagate errors
611                    while let Ok(result) = rx.try_recv() {
612                        result?; // Propagate errors
613                    }
614
615                    // close the stream after all messages are sent, unless it is a standalone stream
616                    if !stream_id.eq(DEFAULT_STREAM_ID){
617                        drop(tx);
618                        while let Some(result) = rx.recv().await {
619                            result?; // Propagate errors
620                        }
621                        return  Ok(());
622                    }
623                }
624                _ = &mut disconnect_rx => {
625                    // Drop tx to close the channel and collect remaining results
626                    drop(tx);
627                    while let Some(result) = rx.recv().await {
628                        result?; // Propagate errors
629                    }
630                                self.remove_transport(stream_id).await?;
631                                // Disconnection detected by keep-alive task
632                                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            // send TaskStatusNotification  if task_store is present and supports subscribe()
728            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        // Task polling for server initiated tasks
738        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}