#![expect(deprecated)]
use std::future::Future;
use rmcp::{
model::{
ClientCapabilities, ClientInfo, CreateMessageRequestParams, CreateMessageResult,
ElicitRequestParams, ElicitResult, ElicitationCapability, FormElicitationCapability,
Implementation, ListRootsResult, LoggingLevel, LoggingMessageNotificationParam,
ProgressNotificationParam, RootsCapabilities, SamplingCapability,
},
service::{NotificationContext, RequestContext, RoleClient},
ClientHandler, ErrorData as McpError,
};
use super::{
elicitation::McpElicitationService, progress::McpProgressRouter, roots::McpRoots,
sampling::McpSamplingService,
};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[allow(clippy::enum_variant_names)]
pub(crate) enum McpServerEvent {
ToolsChanged,
PromptsChanged,
ResourcesChanged,
}
pub(crate) type McpEventSender = tokio::sync::mpsc::UnboundedSender<McpServerEvent>;
pub(crate) type McpEventReceiver = tokio::sync::mpsc::UnboundedReceiver<McpServerEvent>;
pub(crate) struct McpClientServices {
pub(crate) elicit: McpElicitationService,
pub(crate) sample: Option<McpSamplingService>,
}
pub(crate) struct McpClientHandler {
identity: String,
info: ClientInfo,
roots: McpRoots,
progress: McpProgressRouter,
events: McpEventSender,
services: McpClientServices,
}
impl McpClientHandler {
pub(crate) fn new(
identity: impl Into<String>,
roots: McpRoots,
progress: McpProgressRouter,
events: McpEventSender,
services: McpClientServices,
) -> Self {
let identity = identity.into();
Self {
info: client_info(&roots, &services),
identity,
roots,
progress,
events,
services,
}
}
}
fn client_info(roots: &McpRoots, services: &McpClientServices) -> ClientInfo {
let mut capabilities = ClientCapabilities::default();
if !roots.is_empty() {
let mut declared = RootsCapabilities::default();
declared.list_changed = Some(false);
capabilities.roots = Some(declared);
}
if services.elicit.is_available() {
capabilities.elicitation = Some(
ElicitationCapability::new()
.with_form(FormElicitationCapability::new().with_schema_validation(false)),
);
}
if services.sample.is_some() {
capabilities.sampling = Some(SamplingCapability::default());
}
ClientInfo::new(
capabilities,
Implementation::new("rho", env!("CARGO_PKG_VERSION"))
.with_title("Rho")
.with_description("Rho coding agent")
.with_website_url(RHO_WEBSITE),
)
}
const RHO_WEBSITE: &str = "https://github.com/matthewyjiang/rho";
impl ClientHandler for McpClientHandler {
fn get_info(&self) -> ClientInfo {
self.info.clone()
}
fn list_roots(
&self,
_context: RequestContext<RoleClient>,
) -> impl Future<Output = Result<ListRootsResult, McpError>> + Send + '_ {
let roots = self.roots.to_protocol();
std::future::ready(Ok(ListRootsResult::new(roots)))
}
fn create_elicitation(
&self,
request: ElicitRequestParams,
_context: RequestContext<RoleClient>,
) -> impl Future<Output = Result<ElicitResult, McpError>> + Send + '_ {
self.services.elicit.elicit(request)
}
fn create_message(
&self,
params: CreateMessageRequestParams,
_context: RequestContext<RoleClient>,
) -> impl Future<Output = Result<CreateMessageResult, McpError>> + Send + '_ {
self.sample(params)
}
fn on_progress(
&self,
params: ProgressNotificationParam,
_context: NotificationContext<RoleClient>,
) -> impl Future<Output = ()> + Send + '_ {
self.progress.dispatch(params)
}
fn on_logging_message(
&self,
params: LoggingMessageNotificationParam,
_context: NotificationContext<RoleClient>,
) -> impl Future<Output = ()> + Send + '_ {
log_server_message(&self.identity, params);
std::future::ready(())
}
fn on_tool_list_changed(
&self,
_context: NotificationContext<RoleClient>,
) -> impl Future<Output = ()> + Send + '_ {
self.announce(McpServerEvent::ToolsChanged)
}
fn on_prompt_list_changed(
&self,
_context: NotificationContext<RoleClient>,
) -> impl Future<Output = ()> + Send + '_ {
self.announce(McpServerEvent::PromptsChanged)
}
fn on_resource_list_changed(
&self,
_context: NotificationContext<RoleClient>,
) -> impl Future<Output = ()> + Send + '_ {
self.announce(McpServerEvent::ResourcesChanged)
}
}
impl McpClientHandler {
async fn sample(
&self,
params: CreateMessageRequestParams,
) -> Result<CreateMessageResult, McpError> {
let Some(sampling) = self.services.sample.as_ref() else {
return Err(McpError::method_not_found::<
rmcp::model::CreateMessageRequestMethod,
>());
};
sampling.create_message(params).await
}
fn announce(&self, event: McpServerEvent) -> std::future::Ready<()> {
let _ = self.events.send(event);
std::future::ready(())
}
}
fn log_server_message(identity: &str, params: LoggingMessageNotificationParam) {
let logger = params.logger.unwrap_or_else(|| "mcp".into());
let data = match ¶ms.data {
serde_json::Value::String(text) => text.clone(),
other => other.to_string(),
};
match params.level {
LoggingLevel::Debug => {
tracing::debug!(target: "rho::mcp::server", server = %identity, logger = %logger, "{data}");
}
LoggingLevel::Info | LoggingLevel::Notice => {
tracing::info!(target: "rho::mcp::server", server = %identity, logger = %logger, "{data}");
}
LoggingLevel::Warning => {
tracing::warn!(target: "rho::mcp::server", server = %identity, logger = %logger, "{data}");
}
LoggingLevel::Error
| LoggingLevel::Critical
| LoggingLevel::Alert
| LoggingLevel::Emergency => {
tracing::error!(target: "rho::mcp::server", server = %identity, logger = %logger, "{data}");
}
}
}
#[cfg(test)]
#[path = "client_tests.rs"]
mod tests;