1use crate::mcp::protocol::{
6 CallToolParams, CallToolResult, ClientCapabilities, ClientInfo, InitializeParams,
7 InitializeResult, JsonRpcNotification, JsonRpcRequest, ListResourcesResult, ListToolsResult,
8 McpNotification, McpResource, McpTool, ReadResourceParams, ReadResourceResult,
9 ServerCapabilities, PROTOCOL_VERSION,
10};
11use crate::mcp::transport::McpTransport;
12use anyhow::{anyhow, Result};
13use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
14use std::sync::Arc;
15use tokio::sync::RwLock;
16
17pub struct McpClient {
19 pub name: String,
21 transport: Arc<dyn McpTransport>,
23 capabilities: RwLock<ServerCapabilities>,
25 tools: RwLock<Vec<McpTool>>,
27 resources: RwLock<Vec<McpResource>>,
29 request_id: AtomicU64,
31 initialized: AtomicBool,
33}
34
35impl McpClient {
36 pub fn new(name: String, transport: Arc<dyn McpTransport>) -> Self {
38 Self {
39 name,
40 transport,
41 capabilities: RwLock::new(ServerCapabilities::default()),
42 tools: RwLock::new(Vec::new()),
43 resources: RwLock::new(Vec::new()),
44 request_id: AtomicU64::new(1),
45 initialized: AtomicBool::new(false),
46 }
47 }
48
49 fn next_id(&self) -> u64 {
51 self.request_id.fetch_add(1, Ordering::SeqCst)
52 }
53
54 pub async fn initialize(&self) -> Result<InitializeResult> {
56 let params = InitializeParams {
57 protocol_version: PROTOCOL_VERSION.to_string(),
58 capabilities: ClientCapabilities::default(),
59 client_info: ClientInfo {
60 name: "a3s-code".to_string(),
61 version: env!("CARGO_PKG_VERSION").to_string(),
62 },
63 };
64
65 let request = JsonRpcRequest::new(
66 self.next_id(),
67 "initialize",
68 Some(serde_json::to_value(¶ms)?),
69 );
70
71 let response = self.transport.request(request).await?;
72
73 if let Some(error) = response.error {
74 return Err(anyhow!(
75 "MCP initialize error: {} ({})",
76 error.message,
77 error.code
78 ));
79 }
80
81 let result: InitializeResult = serde_json::from_value(
82 response
83 .result
84 .ok_or_else(|| anyhow!("No result in response"))?,
85 )?;
86
87 {
89 let mut caps = self.capabilities.write().await;
90 *caps = result.capabilities.clone();
91 }
92
93 let notification = JsonRpcNotification::new("notifications/initialized", None);
95 self.transport.notify(notification).await?;
96
97 self.initialized.store(true, Ordering::Release);
99
100 tracing::info!(
101 "MCP client '{}' initialized with server '{}' v{}",
102 self.name,
103 result.server_info.name,
104 result.server_info.version
105 );
106
107 Ok(result)
108 }
109
110 pub async fn is_initialized(&self) -> bool {
112 self.initialized.load(Ordering::Acquire)
113 }
114
115 pub fn is_ready(&self) -> bool {
122 self.initialized.load(Ordering::Acquire) && self.transport.is_connected()
123 }
124
125 pub async fn capabilities(&self) -> ServerCapabilities {
127 self.capabilities.read().await.clone()
128 }
129
130 pub async fn list_tools(&self) -> Result<Vec<McpTool>> {
132 let request = JsonRpcRequest::new(self.next_id(), "tools/list", None);
133 let response = self.transport.request(request).await?;
134
135 if let Some(error) = response.error {
136 return Err(anyhow!(
137 "MCP list_tools error: {} ({})",
138 error.message,
139 error.code
140 ));
141 }
142
143 let result: ListToolsResult =
144 serde_json::from_value(response.result.ok_or_else(|| anyhow!("No result"))?)?;
145
146 {
148 let mut tools = self.tools.write().await;
149 *tools = result.tools.clone();
150 }
151
152 Ok(result.tools)
153 }
154
155 pub async fn get_cached_tools(&self) -> Vec<McpTool> {
157 self.tools.read().await.clone()
158 }
159
160 pub async fn call_tool(
162 &self,
163 name: &str,
164 arguments: Option<serde_json::Value>,
165 ) -> Result<CallToolResult> {
166 let params = CallToolParams {
167 name: name.to_string(),
168 arguments,
169 };
170
171 let request = JsonRpcRequest::new(
172 self.next_id(),
173 "tools/call",
174 Some(serde_json::to_value(¶ms)?),
175 );
176
177 let response = self.transport.request(request).await?;
178
179 if let Some(error) = response.error {
180 return Err(anyhow!(
181 "MCP call_tool error: {} ({})",
182 error.message,
183 error.code
184 ));
185 }
186
187 let result: CallToolResult =
188 serde_json::from_value(response.result.ok_or_else(|| anyhow!("No result"))?)?;
189
190 Ok(result)
191 }
192
193 pub async fn list_resources(&self) -> Result<Vec<McpResource>> {
195 let request = JsonRpcRequest::new(self.next_id(), "resources/list", None);
196 let response = self.transport.request(request).await?;
197
198 if let Some(error) = response.error {
199 return Err(anyhow!(
200 "MCP list_resources error: {} ({})",
201 error.message,
202 error.code
203 ));
204 }
205
206 let result: ListResourcesResult =
207 serde_json::from_value(response.result.ok_or_else(|| anyhow!("No result"))?)?;
208
209 {
211 let mut resources = self.resources.write().await;
212 *resources = result.resources.clone();
213 }
214
215 Ok(result.resources)
216 }
217
218 pub async fn read_resource(&self, uri: &str) -> Result<ReadResourceResult> {
220 let params = ReadResourceParams {
221 uri: uri.to_string(),
222 };
223
224 let request = JsonRpcRequest::new(
225 self.next_id(),
226 "resources/read",
227 Some(serde_json::to_value(¶ms)?),
228 );
229
230 let response = self.transport.request(request).await?;
231
232 if let Some(error) = response.error {
233 return Err(anyhow!(
234 "MCP read_resource error: {} ({})",
235 error.message,
236 error.code
237 ));
238 }
239
240 let result: ReadResourceResult =
241 serde_json::from_value(response.result.ok_or_else(|| anyhow!("No result"))?)?;
242
243 Ok(result)
244 }
245
246 pub fn notifications(&self) -> tokio::sync::mpsc::Receiver<McpNotification> {
248 self.transport.notifications()
249 }
250
251 pub async fn close(&self) -> Result<()> {
253 self.initialized.store(false, Ordering::Release);
254 self.transport.close().await
255 }
256
257 pub fn is_connected(&self) -> bool {
259 self.transport.is_connected()
260 }
261}
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266
267 #[test]
268 fn test_client_info() {
269 let info = ClientInfo {
270 name: "test".to_string(),
271 version: "1.0.0".to_string(),
272 };
273 let json = serde_json::to_string(&info).unwrap();
274 assert!(json.contains("test"));
275 }
276
277 #[test]
278 fn test_initialize_params() {
279 let params = InitializeParams {
280 protocol_version: PROTOCOL_VERSION.to_string(),
281 capabilities: ClientCapabilities::default(),
282 client_info: ClientInfo {
283 name: "a3s-code".to_string(),
284 version: "0.1.0".to_string(),
285 },
286 };
287 let json = serde_json::to_string(¶ms).unwrap();
288 assert!(json.contains("protocolVersion"));
289 assert!(json.contains("clientInfo"));
290 }
291
292 #[test]
293 fn test_client_info_serialize() {
294 let info = ClientInfo {
295 name: "test-client".to_string(),
296 version: "2.0.0".to_string(),
297 };
298 let json = serde_json::to_string(&info).unwrap();
299 assert!(json.contains("test-client"));
300 assert!(json.contains("2.0.0"));
301 }
302
303 #[test]
304 fn test_client_info_deserialize() {
305 let json = r#"{"name":"my-client","version":"1.2.3"}"#;
306 let info: ClientInfo = serde_json::from_str(json).unwrap();
307 assert_eq!(info.name, "my-client");
308 assert_eq!(info.version, "1.2.3");
309 }
310
311 #[test]
312 fn test_initialize_params_serialize() {
313 let params = InitializeParams {
314 protocol_version: "2024-11-05".to_string(),
315 capabilities: ClientCapabilities::default(),
316 client_info: ClientInfo {
317 name: "test".to_string(),
318 version: "1.0.0".to_string(),
319 },
320 };
321 let json = serde_json::to_string(¶ms).unwrap();
322 assert!(json.contains("2024-11-05"));
323 assert!(json.contains("capabilities"));
324 }
325
326 #[test]
327 fn test_call_tool_params_serialize() {
328 let params = CallToolParams {
329 name: "test_tool".to_string(),
330 arguments: Some(serde_json::json!({"key": "value"})),
331 };
332 let json = serde_json::to_string(¶ms).unwrap();
333 assert!(json.contains("test_tool"));
334 assert!(json.contains("key"));
335 }
336
337 #[test]
338 fn test_call_tool_params_no_arguments() {
339 let params = CallToolParams {
340 name: "simple_tool".to_string(),
341 arguments: None,
342 };
343 let json = serde_json::to_string(¶ms).unwrap();
344 assert!(json.contains("simple_tool"));
345 }
346
347 #[test]
348 fn test_read_resource_params_serialize() {
349 let params = ReadResourceParams {
350 uri: "file:///test.txt".to_string(),
351 };
352 let json = serde_json::to_string(¶ms).unwrap();
353 assert!(json.contains("file:///test.txt"));
354 }
355
356 #[test]
357 fn test_read_resource_params_deserialize() {
358 let json = r#"{"uri":"http://example.com/resource"}"#;
359 let params: ReadResourceParams = serde_json::from_str(json).unwrap();
360 assert_eq!(params.uri, "http://example.com/resource");
361 }
362
363 #[test]
364 fn test_server_capabilities_default() {
365 let caps = ServerCapabilities::default();
366 let json = serde_json::to_string(&caps).unwrap();
367 assert!(!json.is_empty());
368 }
369
370 #[test]
371 fn test_client_capabilities_default() {
372 let caps = ClientCapabilities::default();
373 let json = serde_json::to_string(&caps).unwrap();
374 assert!(!json.is_empty());
375 }
376
377 #[test]
378 fn test_protocol_version_constant() {
379 assert!(!PROTOCOL_VERSION.is_empty());
380 assert!(PROTOCOL_VERSION.contains("-"));
381 }
382}