use std::future::Future;
use rmcp::RoleServer;
use rmcp::model::{
ClientJsonRpcMessage, ClientRequest, ErrorCode, ErrorData, RequestId, ServerJsonRpcMessage,
};
use rmcp::transport::Transport;
pub(crate) struct TolerantInit<T> {
inner: T,
initialize_seen: bool,
}
impl<T> TolerantInit<T> {
pub(crate) fn new(inner: T) -> Self {
Self {
inner,
initialize_seen: false,
}
}
}
enum PreInit {
Forward,
Reject(ErrorData, RequestId),
Drop,
}
impl<T: Transport<RoleServer>> Transport<RoleServer> for TolerantInit<T> {
type Error = T::Error;
fn send(
&mut self,
item: ServerJsonRpcMessage,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
self.inner.send(item)
}
async fn receive(&mut self) -> Option<ClientJsonRpcMessage> {
loop {
let msg = self.inner.receive().await?;
if self.initialize_seen {
return Some(msg);
}
let decision = match &msg {
ClientJsonRpcMessage::Request(req) => match &req.request {
ClientRequest::InitializeRequest(_) => {
self.initialize_seen = true;
PreInit::Forward
}
ClientRequest::PingRequest(_) => PreInit::Forward,
other => PreInit::Reject(
ErrorData::new(
ErrorCode::METHOD_NOT_FOUND,
other.method().to_string(),
None,
),
req.id.clone(),
),
},
_ => PreInit::Drop,
};
match decision {
PreInit::Forward => return Some(msg),
PreInit::Reject(err, id) => {
let _ = self
.inner
.send(ServerJsonRpcMessage::error(err, Some(id)))
.await;
}
PreInit::Drop => {}
}
}
}
fn close(&mut self) -> impl Future<Output = Result<(), Self::Error>> + Send {
self.inner.close()
}
}
#[cfg(test)]
mod tests {
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use serde_json::json;
use super::*;
struct Mock {
incoming: VecDeque<ClientJsonRpcMessage>,
sent: Arc<Mutex<Vec<ServerJsonRpcMessage>>>,
}
impl Mock {
fn new(msgs: Vec<serde_json::Value>) -> (Self, Arc<Mutex<Vec<ServerJsonRpcMessage>>>) {
let sent = Arc::new(Mutex::new(Vec::new()));
let incoming = msgs
.into_iter()
.map(|v| serde_json::from_value(v).expect("valid client message"))
.collect();
(
Self {
incoming,
sent: sent.clone(),
},
sent,
)
}
}
impl Transport<RoleServer> for Mock {
type Error = std::io::Error;
fn send(
&mut self,
item: ServerJsonRpcMessage,
) -> impl Future<Output = Result<(), Self::Error>> + Send + 'static {
let sent = self.sent.clone();
async move {
sent.lock().unwrap().push(item);
Ok(())
}
}
async fn receive(&mut self) -> Option<ClientJsonRpcMessage> {
self.incoming.pop_front()
}
async fn close(&mut self) -> Result<(), Self::Error> {
Ok(())
}
}
fn discover(id: u32) -> serde_json::Value {
json!({"jsonrpc": "2.0", "id": id, "method": "server/discover", "params": {}})
}
fn initialize(id: u32) -> serde_json::Value {
json!({"jsonrpc": "2.0", "id": id, "method": "initialize", "params": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0"}
}})
}
fn method_of(msg: &ClientJsonRpcMessage) -> Option<&str> {
match msg {
ClientJsonRpcMessage::Request(r) => Some(r.request.method()),
_ => None,
}
}
#[tokio::test]
async fn pre_init_discover_is_rejected_and_initialize_forwarded() {
let (mock, sent) = Mock::new(vec![discover(0), initialize(1)]);
let mut t = TolerantInit::new(mock);
let first = t.receive().await.expect("message");
assert_eq!(method_of(&first), Some("initialize"));
let sent = sent.lock().unwrap();
assert_eq!(sent.len(), 1, "exactly one pre-init reply");
let v = serde_json::to_value(&sent[0]).unwrap();
assert_eq!(v["id"], 0);
assert_eq!(v["error"]["code"], -32601);
assert_eq!(v["error"]["message"], "server/discover");
}
#[tokio::test]
async fn post_init_messages_pass_through_untouched() {
let (mock, sent) = Mock::new(vec![initialize(0), discover(1)]);
let mut t = TolerantInit::new(mock);
assert_eq!(method_of(&t.receive().await.unwrap()), Some("initialize"));
assert_eq!(
method_of(&t.receive().await.unwrap()),
Some("server/discover")
);
assert!(sent.lock().unwrap().is_empty(), "wrapper sent nothing");
}
#[tokio::test]
async fn pre_init_notifications_are_dropped() {
let notif = json!({"jsonrpc": "2.0", "method": "notifications/cancelled", "params": {"requestId": 9}});
let (mock, sent) = Mock::new(vec![notif, initialize(0)]);
let mut t = TolerantInit::new(mock);
assert_eq!(method_of(&t.receive().await.unwrap()), Some("initialize"));
assert!(sent.lock().unwrap().is_empty());
}
#[tokio::test]
async fn pre_init_ping_is_forwarded_for_rmcp_to_answer() {
let ping = json!({"jsonrpc": "2.0", "id": 0, "method": "ping"});
let (mock, sent) = Mock::new(vec![ping, initialize(1)]);
let mut t = TolerantInit::new(mock);
assert_eq!(method_of(&t.receive().await.unwrap()), Some("ping"));
assert_eq!(method_of(&t.receive().await.unwrap()), Some("initialize"));
assert!(sent.lock().unwrap().is_empty());
}
}