1use 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
45pub struct McpClientHandler<S> {
50 client_info: rmcp::model::ClientInfo,
51 sink: S,
52 timeout: Option<Duration>,
54 refresh_timeout: Duration,
56 pub(crate) managed_tools: Arc<RwLock<ManagedToolsState>>,
60 pub(crate) refresh_activity: Arc<Mutex<RefreshActivity>>,
62 next_refresh: Arc<AtomicU64>,
64}
65
66impl<S> McpClientHandler<S>
67where
68 S: ManagedToolSink + Send + Sync + 'static,
69{
70 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 #[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 #[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 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 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 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 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}