Skip to main content

rig_rmcp/
handler.rs

1//! Synchronizes managed tool registrations with an MCP server's tool list.
2//!
3//! ```
4//! use rig_core::tool::ManagedToolSink;
5//! use rig_rmcp::{McpClientHandler, rmcp};
6//!
7//! fn handler<S: ManagedToolSink + Send + Sync + 'static>(sink: S) -> McpClientHandler<S> {
8//!     McpClientHandler::new(rmcp::model::ClientInfo::default(), sink)
9//! }
10//! ```
11
12use std::collections::HashMap;
13use std::sync::{
14    Arc,
15    atomic::{AtomicU64, Ordering},
16};
17use std::time::Duration;
18
19use rmcp::ServiceExt;
20use rmcp::model::{ClientRequest, ListToolsRequest, PaginatedRequestParams, ServerResult};
21use tokio::sync::{Mutex, RwLock};
22
23use rig_core::message::EmptyToolName;
24use rig_core::tool::{DynamicTool, ManagedToolSink, ManagedToolToken};
25
26use crate::{
27    DEFAULT_MCP_REFRESH_TIMEOUT, DEFAULT_MCP_TOOL_TIMEOUT, McpClientError, McpTool,
28    send_mcp_request,
29};
30
31#[derive(Default)]
32pub(crate) struct ManagedToolsState {
33    pub(crate) registrations: HashMap<String, ManagedToolToken>,
34    pub(crate) committed_refresh: u64,
35}
36
37#[derive(Default)]
38pub(crate) struct RefreshActivity {
39    pub(crate) active: usize,
40    pub(crate) dirty: bool,
41}
42
43pub(crate) const MAX_CONCURRENT_REFRESHES: usize = 2;
44
45/// Registers MCP tools in a [`ManagedToolSink`] and refreshes them on list changes.
46/// Refreshes replace or remove only registrations still owned by this handler,
47/// preserving newer same-name registrations. [`Self::connect`] performs the
48/// initial fetch and returns the service that keeps the connection alive.
49pub struct McpClientHandler<S> {
50    client_info: rmcp::model::ClientInfo,
51    sink: S,
52    /// Per-call timeout for registered tools, initially [`DEFAULT_MCP_TOOL_TIMEOUT`].
53    timeout: Option<Duration>,
54    /// Deadline for initial and list-changed tool-list fetches.
55    refresh_timeout: Duration,
56    /// Tracks the exact registry generation installed for each tool. Refreshes
57    /// only mutate a name while this generation remains current, so a newer
58    /// local or peer-handler registration cannot be deleted or overwritten.
59    pub(crate) managed_tools: Arc<RwLock<ManagedToolsState>>,
60    /// Bounds notification-driven list fetches and coalesces excess signals.
61    pub(crate) refresh_activity: Arc<Mutex<RefreshActivity>>,
62    /// Monotonic identity assigned when each tool-list fetch begins.
63    next_refresh: Arc<AtomicU64>,
64}
65
66impl<S> McpClientHandler<S>
67where
68    S: ManagedToolSink + Send + Sync + 'static,
69{
70    /// Create a new handler with the given client info and tool sink.
71    ///
72    /// With rig-agent, pass a clone of the agent's `ToolServerHandle` so tool
73    /// updates are reflected in agent requests. Registered tools get
74    /// [`DEFAULT_MCP_TOOL_TIMEOUT`]; change it with [`McpClientHandler::with_timeout`].
75    pub fn new(client_info: rmcp::model::ClientInfo, sink: S) -> Self {
76        Self {
77            client_info,
78            sink,
79            timeout: Some(DEFAULT_MCP_TOOL_TIMEOUT),
80            refresh_timeout: DEFAULT_MCP_REFRESH_TIMEOUT,
81            managed_tools: Arc::new(RwLock::new(ManagedToolsState::default())),
82            refresh_activity: Arc::new(Mutex::new(RefreshActivity::default())),
83            next_refresh: Arc::new(AtomicU64::new(0)),
84        }
85    }
86
87    /// Set (or clear) the per-call timeout applied to every MCP tool this handler
88    /// registers. Pass a [`Duration`] to bound calls, or `None` to disable.
89    ///
90    /// This applies the same setting to every tool managed by the handler.
91    #[must_use = "the setting applies to the returned value"]
92    pub fn with_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
93        self.timeout = timeout.into();
94        self
95    }
96
97    /// Set the deadline for initial and list-changed tool-list fetches.
98    #[must_use = "the setting applies to the returned value"]
99    pub fn with_refresh_timeout(mut self, timeout: Duration) -> Self {
100        self.refresh_timeout = timeout;
101        self
102    }
103
104    /// Build the dynamic tool with this handler's configured timeout.
105    pub(crate) fn build_tool(
106        &self,
107        tool: rmcp::model::Tool,
108        client: rmcp::service::ServerSink,
109    ) -> Result<DynamicTool, EmptyToolName> {
110        McpTool::from_mcp_server(tool, client)
111            .with_timeout(self.timeout)
112            .try_into()
113    }
114
115    pub(crate) fn begin_refresh(&self) -> u64 {
116        self.next_refresh.fetch_add(1, Ordering::SeqCst) + 1
117    }
118
119    pub(crate) async fn fetch_tools(
120        &self,
121        peer: &rmcp::service::ServerSink,
122    ) -> Result<Vec<DynamicTool>, McpClientError> {
123        let deadline = tokio::time::Instant::now() + self.refresh_timeout;
124        let mut tools = Vec::new();
125        let mut cursor = None;
126
127        loop {
128            let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
129            if remaining.is_zero() {
130                return Err(McpClientError::ToolFetchTimeout(self.refresh_timeout));
131            }
132            let mut params = PaginatedRequestParams::default();
133            params.cursor = cursor;
134            let response = send_mcp_request(
135                peer,
136                ClientRequest::ListToolsRequest(ListToolsRequest::with_param(params)),
137                Some((deadline, self.refresh_timeout)),
138            )
139            .await
140            .map_err(|error| match error {
141                rmcp::ServiceError::Timeout { .. } => {
142                    McpClientError::ToolFetchTimeout(self.refresh_timeout)
143                }
144                error => McpClientError::ToolFetch(error),
145            })?;
146            let ServerResult::ListToolsResult(page) = response else {
147                return Err(McpClientError::ToolFetch(
148                    rmcp::ServiceError::UnexpectedResponse,
149                ));
150            };
151            tools.extend(page.tools);
152            cursor = page.next_cursor;
153            if cursor.is_none() {
154                break;
155            }
156        }
157
158        // One nameless tool is a server bug; it should not hide the server's other tools.
159        Ok(tools
160            .into_iter()
161            .filter_map(|tool| match self.build_tool(tool, peer.clone()) {
162                Ok(tool) => Some(tool),
163                Err(error) => {
164                    tracing::warn!(%error, "skipping an MCP tool the server listed without a name");
165                    None
166                }
167            })
168            .collect())
169    }
170
171    pub(crate) async fn try_start_refresh(&self) -> bool {
172        let mut activity = self.refresh_activity.lock().await;
173        if activity.active >= MAX_CONCURRENT_REFRESHES {
174            activity.dirty = true;
175            false
176        } else {
177            activity.active += 1;
178            true
179        }
180    }
181
182    pub(crate) async fn finish_or_restart_refresh(&self) -> bool {
183        let mut activity = self.refresh_activity.lock().await;
184        if activity.dirty {
185            activity.dirty = false;
186            true
187        } else {
188            activity.active -= 1;
189            false
190        }
191    }
192
193    pub(crate) async fn commit_initial(&self, refresh: u64, tools: Vec<DynamicTool>) {
194        let mut managed = self.managed_tools.write().await;
195        if refresh <= managed.committed_refresh {
196            tracing::debug!(refresh, "discarding stale initial MCP tool list");
197            return;
198        }
199        managed.registrations = self.sink.add_managed_tools(tools);
200        managed.committed_refresh = refresh;
201    }
202
203    pub(crate) async fn commit_refresh(&self, refresh: u64, tools: Vec<DynamicTool>) -> bool {
204        let mut managed = self.managed_tools.write().await;
205        if refresh <= managed.committed_refresh {
206            tracing::debug!(refresh, "discarding stale MCP tool-list response");
207            return false;
208        }
209        let expected = managed.registrations.clone();
210        managed.registrations = self.sink.reconcile_managed_tools(expected, tools);
211        managed.committed_refresh = refresh;
212        true
213    }
214
215    /// Connect to an MCP server, fetch the initial tool list, and register
216    /// all tools with the tool server.
217    ///
218    /// Returns the running MCP service. The connection stays alive as long as the
219    /// returned `RunningService` is held. When the server sends
220    /// `notifications/tools/list_changed`, this handler automatically re-fetches
221    /// and re-registers tools into the sink.
222    ///
223    /// # Errors
224    ///
225    /// Returns [`McpClientError`] if the connection or initial tool fetch fails.
226    pub async fn connect<T, E, A>(
227        self,
228        transport: T,
229    ) -> Result<rmcp::service::RunningService<rmcp::service::RoleClient, Self>, McpClientError>
230    where
231        T: rmcp::transport::IntoTransport<rmcp::service::RoleClient, E, A>,
232        E: std::error::Error + Send + Sync + 'static,
233    {
234        let service = ServiceExt::serve(self, transport).await?;
235
236        let handler = service.service();
237        let refresh = handler.begin_refresh();
238        let tools = handler.fetch_tools(service.peer()).await?;
239        handler.commit_initial(refresh, tools).await;
240
241        Ok(service)
242    }
243}
244
245impl<S> rmcp::handler::client::ClientHandler for McpClientHandler<S>
246where
247    S: ManagedToolSink + Send + Sync + 'static,
248{
249    fn get_info(&self) -> rmcp::model::ClientInfo {
250        self.client_info.clone()
251    }
252
253    async fn on_tool_list_changed(
254        &self,
255        context: rmcp::service::NotificationContext<rmcp::service::RoleClient>,
256    ) {
257        if !self.try_start_refresh().await {
258            return;
259        }
260
261        loop {
262            let refresh = self.begin_refresh();
263            // Network IO is deliberately outside the ownership lock. Up to two
264            // fetches may overlap so a newer snapshot can bypass one stalled
265            // request; further notifications coalesce into one follow-up fetch.
266            match self.fetch_tools(&context.peer).await {
267                Ok(tools) => {
268                    if self.commit_refresh(refresh, tools).await {
269                        let tool_count = self.managed_tools.read().await.registrations.len();
270                        tracing::info!(tool_count, "MCP tool list refreshed successfully");
271                    }
272                }
273                Err(error) => tracing::error!("Failed to re-fetch MCP tool list: {error}"),
274            }
275
276            if !self.finish_or_restart_refresh().await {
277                break;
278            }
279        }
280    }
281}