Skip to main content

a3s_code_core/mcp/
client.rs

1//! MCP Client
2//!
3//! Provides a high-level client for interacting with MCP servers.
4
5use 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
17/// MCP client for communicating with MCP servers
18pub struct McpClient {
19    /// Server name
20    pub name: String,
21    /// Transport layer
22    transport: Arc<dyn McpTransport>,
23    /// Server capabilities (after initialization)
24    capabilities: RwLock<ServerCapabilities>,
25    /// Cached tools
26    tools: RwLock<Vec<McpTool>>,
27    /// Cached resources
28    resources: RwLock<Vec<McpResource>>,
29    /// Request ID counter
30    request_id: AtomicU64,
31    /// Initialized flag
32    initialized: AtomicBool,
33}
34
35impl McpClient {
36    /// Create a new MCP client with the given transport
37    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    /// Get next request ID
50    fn next_id(&self) -> u64 {
51        self.request_id.fetch_add(1, Ordering::SeqCst)
52    }
53
54    /// Initialize the MCP connection
55    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(&params)?),
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        // Store capabilities
88        {
89            let mut caps = self.capabilities.write().await;
90            *caps = result.capabilities.clone();
91        }
92
93        // Send initialized notification
94        let notification = JsonRpcNotification::new("notifications/initialized", None);
95        self.transport.notify(notification).await?;
96
97        // Mark as initialized
98        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    /// Check if client is initialized
111    pub async fn is_initialized(&self) -> bool {
112        self.initialized.load(Ordering::Acquire)
113    }
114
115    /// Return whether this exact client completed MCP initialization and its
116    /// transport is still connected.
117    ///
118    /// This synchronous readiness check is used while freezing a capability
119    /// Run. It never performs discovery or reconnects through mutable manager
120    /// state.
121    pub fn is_ready(&self) -> bool {
122        self.initialized.load(Ordering::Acquire) && self.transport.is_connected()
123    }
124
125    /// Get server capabilities
126    pub async fn capabilities(&self) -> ServerCapabilities {
127        self.capabilities.read().await.clone()
128    }
129
130    /// List available tools
131    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        // Cache tools
147        {
148            let mut tools = self.tools.write().await;
149            *tools = result.tools.clone();
150        }
151
152        Ok(result.tools)
153    }
154
155    /// Get cached tools
156    pub async fn get_cached_tools(&self) -> Vec<McpTool> {
157        self.tools.read().await.clone()
158    }
159
160    /// Call a tool
161    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(&params)?),
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    /// List available resources
194    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        // Cache resources
210        {
211            let mut resources = self.resources.write().await;
212            *resources = result.resources.clone();
213        }
214
215        Ok(result.resources)
216    }
217
218    /// Read a resource
219    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(&params)?),
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    /// Get notification receiver
247    pub fn notifications(&self) -> tokio::sync::mpsc::Receiver<McpNotification> {
248        self.transport.notifications()
249    }
250
251    /// Close the client
252    pub async fn close(&self) -> Result<()> {
253        self.initialized.store(false, Ordering::Release);
254        self.transport.close().await
255    }
256
257    /// Check if connected
258    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(&params).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(&params).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(&params).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(&params).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(&params).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}