use std::{net::SocketAddr, sync::Arc};
use connectrpc::Router;
use rmcp::{
ServerHandler,
transport::streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService, session::never::NeverSessionManager,
},
};
use tokio_util::sync::CancellationToken;
pub fn build_router<H>(path: &str, handler: H) -> axum::Router
where
H: ServerHandler + Clone + Send + Sync + 'static,
{
let mut config = StreamableHttpServerConfig::default();
config.legacy_session_mode = false;
config.sse_keep_alive = Some(std::time::Duration::from_secs(15));
let config = config.disable_allowed_hosts().disable_allowed_origins();
let service = StreamableHttpService::new(
move || Ok(handler.clone()),
Arc::new(NeverSessionManager::default()),
config,
);
let (health, _checker) = connectrpc_health::install_static(Router::new(), [] as [&str; 0]);
axum::Router::new()
.nest_service(path, service)
.fallback_service(health.into_axum_service())
}
pub async fn serve<H>(bind: SocketAddr, path: &str, handler: H) -> anyhow::Result<()>
where
H: ServerHandler + Clone + Send + Sync + 'static,
{
let router = build_router(path, handler);
let listener = tokio::net::TcpListener::bind(bind).await?;
let local = listener.local_addr()?;
tracing::info!(addr = %local, path, "MCP server listening");
println!("MCP server listening on http://{local}{path}");
let ct = CancellationToken::new();
let server_ct = ct.clone();
let server = tokio::spawn(async move {
let _ = axum::serve(listener, router)
.with_graceful_shutdown(async move { server_ct.cancelled_owned().await })
.await;
});
tokio::signal::ctrl_c().await?;
tracing::info!("shutting down");
ct.cancel();
let _ = server.await;
Ok(())
}
#[cfg(test)]
mod tests {
#![allow(clippy::pedantic, clippy::nursery, missing_docs)]
use axum::body::Body;
use axum::http::{Request, StatusCode, header};
use rmcp::model::{Implementation, InitializeResult, ServerCapabilities, ServerInfo};
use tower::ServiceExt;
use super::*;
#[derive(Clone)]
struct NoopServer;
impl ServerHandler for NoopServer {
fn get_info(&self) -> ServerInfo {
InitializeResult::new(ServerCapabilities::builder().build())
.with_server_info(Implementation::new("noop", "0.0.0"))
}
}
#[tokio::test]
async fn health_check_reports_serving() {
let router = build_router("/mcp", NoopServer);
let resp = router
.oneshot(
Request::post("/grpc.health.v1.Health/Check")
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from("{}"))
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
let value: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(value["status"], "SERVING");
}
}