fix: check server capability when client sends requests (#558)

This commit is contained in:
Salman Mohammed
2025-01-08 16:42:23 -05:00
committed by GitHub
parent bb706fb138
commit 640c38b736
6 changed files with 80 additions and 11 deletions
+1 -1
View File
@@ -33,7 +33,7 @@ impl Capabilities {
/// Add a new MCP system based on the provided client type
// TODO IMPORTANT need to ensure this times out if the system command is broken!
pub async fn add_system(&mut self, config: SystemConfig) -> SystemResult<()> {
let client: McpClient = match config {
let mut client: McpClient = match config {
SystemConfig::Sse { ref uri } => {
let transport = SseTransport::new(uri);
McpClient::new(transport.start().await?)
+7 -3
View File
@@ -58,7 +58,6 @@ impl DefaultAgent {
&resources,
Some(model_name),
);
let mut status_content: Vec<String> = Vec::new();
if approx_count > target_limit {
@@ -217,6 +216,7 @@ impl Agent for DefaultAgent {
}
// Update conversation history for the start of the reply
let resources = capabilities.get_resources().await?;
let mut messages = self
.prepare_inference(
&system_prompt,
@@ -224,8 +224,12 @@ impl Agent for DefaultAgent {
messages,
&Vec::new(),
estimated_limit,
&capabilities.provider().get_model_config().model_name,
&capabilities.get_resources().await?,
&capabilities
.provider()
.get_model_config()
.model_name
.clone(),
&resources,
)
.await?;
+1 -1
View File
@@ -22,7 +22,7 @@ async fn main() -> Result<()> {
let handle = transport.start().await?;
// Create client
let client = McpClient::new(handle);
let mut client = McpClient::new(handle);
println!("Client created\n");
// Initialize
+5 -1
View File
@@ -21,7 +21,7 @@ async fn main() -> Result<(), ClientError> {
let transport_handle = transport.start().await?;
// 3) Create the client
let client = McpClient::new(transport_handle);
let mut client = McpClient::new(transport_handle);
// Initialize
let server_info = client
@@ -45,5 +45,9 @@ async fn main() -> Result<(), ClientError> {
.await?;
println!("Tool result: {tool_result:?}\n");
// List resources
let resources = client.list_resources().await?;
println!("Available resources: {resources:?}\n");
Ok(())
}
@@ -29,7 +29,7 @@ async fn main() -> Result<(), ClientError> {
let transport_handle = transport.start().await.unwrap();
// Create client
let client = McpClient::new(transport_handle);
let mut client = McpClient::new(transport_handle);
// Initialize
let server_info = client
+65 -4
View File
@@ -1,16 +1,16 @@
use std::sync::atomic::{AtomicU64, Ordering};
use crate::transport::TransportHandle;
use mcp_core::protocol::{
CallToolResult, InitializeResult, JsonRpcError, JsonRpcMessage, JsonRpcNotification,
JsonRpcRequest, JsonRpcResponse, ListResourcesResult, ListToolsResult, ReadResourceResult,
ServerCapabilities, METHOD_NOT_FOUND,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use thiserror::Error;
use tokio::sync::Mutex;
use tower::{Service, ServiceExt};
use crate::transport::TransportHandle; // for Service::ready()
use tower::{Service, ServiceExt}; // for Service::ready()
/// Error type for MCP client operations.
#[derive(Debug, Error)]
@@ -27,6 +27,9 @@ pub enum Error {
#[error("Unexpected response from server")]
UnexpectedResponse,
#[error("Not initialized")]
NotInitialized,
#[error("Timeout or service not ready")]
NotReady,
}
@@ -55,6 +58,7 @@ pub struct InitializeParams {
pub struct McpClient {
service: Mutex<TransportHandle>,
next_id: AtomicU64,
server_capabilities: Option<ServerCapabilities>,
}
impl McpClient {
@@ -63,6 +67,7 @@ impl McpClient {
Self {
service: Mutex::new(transport_handle),
next_id: AtomicU64::new(1),
server_capabilities: None, // set during initialization
}
}
@@ -135,7 +140,7 @@ impl McpClient {
}
pub async fn initialize(
&self,
&mut self,
info: ClientInfo,
capabilities: ClientCapabilities,
) -> Result<InitializeResult, Error> {
@@ -151,24 +156,80 @@ impl McpClient {
self.send_notification("notifications/initialized", serde_json::json!({}))
.await?;
self.server_capabilities = Some(result.capabilities.clone());
Ok(result)
}
fn completed_initialization(&self) -> bool {
self.server_capabilities.is_some()
}
pub async fn list_resources(&self) -> Result<ListResourcesResult, Error> {
if !self.completed_initialization() {
return Err(Error::NotInitialized);
}
// If resources is not supported, return an empty list
if self
.server_capabilities
.as_ref()
.unwrap()
.resources
.is_none()
{
return Ok(ListResourcesResult { resources: vec![] });
}
self.send_request("resources/list", serde_json::json!({}))
.await
}
pub async fn read_resource(&self, uri: &str) -> Result<ReadResourceResult, Error> {
if !self.completed_initialization() {
return Err(Error::NotInitialized);
}
// If resources is not supported, return an error
if self
.server_capabilities
.as_ref()
.unwrap()
.resources
.is_none()
{
return Err(Error::RpcError {
code: METHOD_NOT_FOUND,
message: "Server does not support 'resources' capability".to_string(),
});
}
let params = serde_json::json!({ "uri": uri });
self.send_request("resources/read", params).await
}
pub async fn list_tools(&self) -> Result<ListToolsResult, Error> {
if !self.completed_initialization() {
return Err(Error::NotInitialized);
}
// If tools is not supported, return an empty list
if self.server_capabilities.as_ref().unwrap().tools.is_none() {
return Ok(ListToolsResult { tools: vec![] });
}
self.send_request("tools/list", serde_json::json!({})).await
}
pub async fn call_tool(&self, name: &str, arguments: Value) -> Result<CallToolResult, Error> {
if !self.completed_initialization() {
return Err(Error::NotInitialized);
}
// If tools is not supported, return an error
if self.server_capabilities.as_ref().unwrap().tools.is_none() {
return Err(Error::RpcError {
code: METHOD_NOT_FOUND,
message: "Server does not support 'tools' capability".to_string(),
});
}
let params = serde_json::json!({ "name": name, "arguments": arguments });
self.send_request("tools/call", params).await
}