use crate::builder::Server;
use crate::context::{CancellationToken, Context, ContextData, Peer};
use crate::dispatch::{PromptSlot, ResourceSlot, TaskSlot, ToolSlot};
use crate::handler::ServerHandler;
use crate::router::{route_prompts, route_resources, route_tasks, route_tools};
use futures::channel::{mpsc, oneshot};
use mcpkit_core::capability::{ClientCapabilities, ServerCapabilities};
use mcpkit_core::error::McpError;
use mcpkit_core::protocol::{Message, Notification, ProgressToken, Request, RequestId, Response};
use mcpkit_core::protocol_version::ProtocolVersion;
use mcpkit_transport::Transport;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::RwLock;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
pub struct ServerState {
pub client_caps: RwLock<ClientCapabilities>,
pub server_caps: ServerCapabilities,
pub initialized: AtomicBool,
pub cancellations: RwLock<HashMap<String, CancellationToken>>,
pub negotiated_version: RwLock<Option<ProtocolVersion>>,
outbound: crate::adapter_peer::SessionOutbound,
ambient_tx: mpsc::UnboundedSender<Notification>,
ambient_rx: std::sync::Mutex<Option<mpsc::UnboundedReceiver<Notification>>>,
}
impl ServerState {
#[must_use]
pub fn new(server_caps: ServerCapabilities) -> Self {
let (ambient_tx, ambient_rx) = mpsc::unbounded();
Self {
client_caps: RwLock::new(ClientCapabilities::default()),
server_caps,
initialized: AtomicBool::new(false),
cancellations: RwLock::new(HashMap::new()),
negotiated_version: RwLock::new(None),
outbound: crate::adapter_peer::SessionOutbound::new(),
ambient_tx,
ambient_rx: std::sync::Mutex::new(Some(ambient_rx)),
}
}
pub fn publish_notification(&self, notification: Notification) {
let _ = self.ambient_tx.unbounded_send(notification);
}
fn take_ambient_receiver(&self) -> Option<mpsc::UnboundedReceiver<Notification>> {
self.ambient_rx.lock().ok()?.take()
}
pub(crate) fn next_outbound_id(&self) -> RequestId {
self.outbound.next_id()
}
pub(crate) fn register_outbound(&self, id: RequestId) -> oneshot::Receiver<Response> {
self.outbound.register(id)
}
pub(crate) fn remove_outbound(&self, id: &RequestId) {
self.outbound.remove(id);
}
pub(crate) fn route_response(&self, response: Response) {
let id = response.id.clone();
if !self.outbound.resolve(response) {
tracing::debug!(id = %id, "response did not match a pending request");
}
}
pub(crate) fn fail_pending_requests(&self) {
self.outbound.fail_all();
}
pub fn protocol_version(&self) -> Option<ProtocolVersion> {
self.negotiated_version.read().ok().and_then(|guard| *guard)
}
pub fn set_protocol_version(&self, version: ProtocolVersion) {
if let Ok(mut guard) = self.negotiated_version.write() {
*guard = Some(version);
}
}
pub fn client_caps(&self) -> ClientCapabilities {
self.client_caps
.read()
.map(|guard| guard.clone())
.unwrap_or_default()
}
pub fn set_client_caps(&self, caps: ClientCapabilities) {
if let Ok(mut guard) = self.client_caps.write() {
*guard = caps;
}
}
pub fn is_initialized(&self) -> bool {
self.initialized.load(Ordering::Acquire)
}
pub fn set_initialized(&self) {
self.initialized.store(true, Ordering::Release);
}
pub fn register_cancellation(&self, request_id: &str, token: CancellationToken) {
if let Ok(mut cancellations) = self.cancellations.write() {
cancellations.insert(request_id.to_string(), token);
}
}
pub fn cancel_request(&self, request_id: &str) {
if let Ok(cancellations) = self.cancellations.read() {
if let Some(token) = cancellations.get(request_id) {
token.cancel();
}
}
}
pub fn remove_cancellation(&self, request_id: &str) {
if let Ok(mut cancellations) = self.cancellations.write() {
cancellations.remove(request_id);
}
}
}
#[derive(Clone)]
struct OutboundCtx {
state: Arc<ServerState>,
timeout: Duration,
}
pub struct TransportPeer<T: Transport> {
transport: Arc<T>,
outbound: Option<OutboundCtx>,
}
impl<T: Transport> TransportPeer<T> {
pub const fn new(transport: Arc<T>) -> Self {
Self {
transport,
outbound: None,
}
}
pub(crate) fn with_outbound(
transport: Arc<T>,
state: Arc<ServerState>,
timeout: Duration,
) -> Self {
Self {
transport,
outbound: Some(OutboundCtx { state, timeout }),
}
}
}
impl<T: Transport + 'static> Peer for TransportPeer<T>
where
T::Error: Into<McpError>,
{
fn notify(
&self,
notification: Notification,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<(), McpError>> + Send + '_>>
{
let transport = self.transport.clone();
Box::pin(async move {
transport
.send(Message::Notification(notification))
.await
.map_err(std::convert::Into::into)
})
}
fn request(
&self,
method: std::borrow::Cow<'static, str>,
params: Option<serde_json::Value>,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<Response, McpError>> + Send + '_>>
{
let Some(outbound) = self.outbound.clone() else {
return Box::pin(async {
Err(McpError::internal(
"this peer does not support server-initiated requests",
))
});
};
let transport = self.transport.clone();
Box::pin(async move {
use futures::future::{Either, select};
let id = outbound.state.next_outbound_id();
let rx = outbound.state.register_outbound(id.clone());
let request = match params {
Some(p) => Request::with_params(method, id.clone(), p),
None => Request::new(method, id.clone()),
};
transport
.send(Message::Request(request))
.await
.map_err(std::convert::Into::into)?;
let sleep = mcpkit_transport::runtime::sleep(outbound.timeout);
futures::pin_mut!(sleep);
match select(rx, sleep).await {
Either::Left((Ok(response), _)) => Ok(response),
Either::Left((Err(_canceled), _)) => {
outbound.state.remove_outbound(&id);
Err(McpError::internal(
"response channel closed before a reply arrived",
))
}
Either::Right(((), _)) => {
outbound.state.remove_outbound(&id);
Err(McpError::internal(format!(
"server-initiated request timed out after {:?}",
outbound.timeout
)))
}
}
})
}
}
#[derive(Clone)]
pub struct ServerNotifier {
peer: Arc<dyn Peer>,
}
impl ServerNotifier {
pub async fn notify(
&self,
method: impl Into<std::borrow::Cow<'static, str>>,
params: Option<serde_json::Value>,
) -> Result<(), McpError> {
let notification = match params {
Some(p) => Notification::with_params(method, p),
None => Notification::new(method),
};
self.peer.notify(notification).await
}
pub async fn log(
&self,
level: mcpkit_core::types::LoggingLevel,
logger: Option<&str>,
data: serde_json::Value,
) -> Result<(), McpError> {
let params = mcpkit_core::types::LoggingMessageNotificationParams {
logger: logger.map(String::from),
..mcpkit_core::types::LoggingMessageNotificationParams::new(level, data)
};
self.notify(
crate::router::notifications::MESSAGE,
Some(serde_json::to_value(params)?),
)
.await
}
pub async fn tools_list_changed(&self) -> Result<(), McpError> {
self.notify(crate::router::notifications::TOOLS_LIST_CHANGED, None)
.await
}
pub async fn resources_list_changed(&self) -> Result<(), McpError> {
self.notify(crate::router::notifications::RESOURCES_LIST_CHANGED, None)
.await
}
pub async fn prompts_list_changed(&self) -> Result<(), McpError> {
self.notify(crate::router::notifications::PROMPTS_LIST_CHANGED, None)
.await
}
pub async fn resource_updated(&self, uri: impl Into<String>) -> Result<(), McpError> {
self.notify(
crate::router::notifications::RESOURCES_UPDATED,
Some(serde_json::json!({ "uri": uri.into() })),
)
.await
}
pub async fn elicitation_complete(
&self,
elicitation_id: impl Into<String>,
) -> Result<(), McpError> {
self.notify(
crate::router::notifications::ELICITATION_COMPLETE,
Some(serde_json::json!({ "elicitationId": elicitation_id.into() })),
)
.await
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct RuntimeConfig {
pub auto_initialized: bool,
pub max_concurrent_requests: usize,
pub outbound_request_timeout: Duration,
pub default_task_ttl_ms: Option<u64>,
pub default_task_poll_interval_ms: Option<u64>,
pub task_status_notifications: bool,
}
impl Default for RuntimeConfig {
fn default() -> Self {
Self {
auto_initialized: true,
max_concurrent_requests: 100,
outbound_request_timeout: Duration::from_secs(60),
default_task_ttl_ms: Some(crate::capability::tasks::DEFAULT_TASK_TTL_MS),
default_task_poll_interval_ms: None,
task_status_notifications: true,
}
}
}
impl RuntimeConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub const fn auto_initialized(mut self, yes: bool) -> Self {
self.auto_initialized = yes;
self
}
#[must_use]
pub const fn max_concurrent_requests(mut self, max: usize) -> Self {
self.max_concurrent_requests = max;
self
}
#[must_use]
pub const fn outbound_request_timeout(mut self, timeout: Duration) -> Self {
self.outbound_request_timeout = timeout;
self
}
#[must_use]
pub const fn default_task_ttl_ms(mut self, ttl_ms: Option<u64>) -> Self {
self.default_task_ttl_ms = ttl_ms;
self
}
#[must_use]
pub const fn default_task_poll_interval_ms(mut self, poll_interval_ms: Option<u64>) -> Self {
self.default_task_poll_interval_ms = poll_interval_ms;
self
}
#[must_use]
pub const fn task_status_notifications(mut self, yes: bool) -> Self {
self.task_status_notifications = yes;
self
}
}
pub struct ServerRuntime<S, Tr>
where
Tr: Transport,
{
server: S,
transport: Arc<Tr>,
state: Arc<ServerState>,
task_store: Arc<crate::capability::tasks::TaskManager>,
config: RuntimeConfig,
}
struct BackgroundExec {
handle: crate::capability::tasks::TaskHandle,
name: String,
args: mcpkit_core::types::Object,
ctx_data: ContextData,
cancel: CancellationToken,
}
enum TaskBegin {
NotApplicable,
Rejected,
Deferred(Box<BackgroundExec>),
}
async fn drive_sets<F1, F2, F3>(
in_flight: &mut futures::stream::FuturesUnordered<F1>,
background: &mut futures::stream::FuturesUnordered<F2>,
notifications: &mut futures::stream::FuturesUnordered<F3>,
) -> Option<BackgroundExec>
where
F1: std::future::Future<Output = Option<BackgroundExec>>,
F2: std::future::Future<Output = ()>,
F3: std::future::Future<Output = Result<(), McpError>>,
{
use futures::future::{Either, select};
use futures::stream::StreamExt;
use std::future::pending;
use std::pin::pin;
let requests = pin!(async {
if in_flight.is_empty() {
pending::<Option<BackgroundExec>>().await
} else {
in_flight.next().await.flatten()
}
});
let tasks = pin!(async {
if background.is_empty() {
pending::<()>().await;
} else {
background.next().await;
}
});
let notifs = pin!(async {
if notifications.is_empty() {
pending::<()>().await;
} else if let Some(Err(e)) = notifications.next().await {
tracing::error!(error = %e, "Error handling notification");
}
});
let unit = pin!(async {
let _ = select(tasks, notifs).await;
});
match select(requests, unit).await {
Either::Left((res, _)) => res,
Either::Right(((), _)) => None,
}
}
impl<S, Tr> ServerRuntime<S, Tr>
where
S: RequestRouter + Send + Sync,
Tr: Transport + 'static,
Tr::Error: Into<McpError>,
{
pub const fn state(&self) -> &Arc<ServerState> {
&self.state
}
#[must_use]
pub fn notifier(&self) -> ServerNotifier {
ServerNotifier {
peer: Arc::new(TransportPeer::new(self.transport.clone())),
}
}
pub async fn run(&self) -> Result<(), McpError> {
use futures::future::{Either, select};
use futures::stream::{FuturesUnordered, StreamExt};
enum Step {
Message(Option<Message>),
Progress(Option<Box<BackgroundExec>>),
Ambient(Notification),
}
async fn next_ambient(
slot: &mut Option<mpsc::UnboundedReceiver<Notification>>,
) -> Notification {
loop {
match slot {
Some(rx) => match rx.next().await {
Some(notification) => return notification,
None => *slot = None,
},
None => std::future::pending::<()>().await,
}
}
}
let mut ambient = self.state.take_ambient_receiver();
let max = self.config.max_concurrent_requests.max(1);
let mut in_flight = FuturesUnordered::new();
let mut background = FuturesUnordered::new();
let mut notifications = FuturesUnordered::new();
let mut queued: std::collections::VecDeque<Request> = std::collections::VecDeque::new();
let outcome = loop {
while in_flight.len() < max {
let Some(request) = queued.pop_front() else {
break;
};
in_flight.push(self.handle_request_isolated(request));
}
let recv = std::pin::pin!(self.transport.recv());
let published = std::pin::pin!(next_ambient(&mut ambient));
let idle = in_flight.is_empty() && background.is_empty() && notifications.is_empty();
let step = if idle {
match select(recv, published).await {
Either::Left((Ok(opt), _)) => Step::Message(opt),
Either::Left((Err(e), _)) => break Err(e.into()),
Either::Right((notification, _)) => Step::Ambient(notification),
}
} else {
let progress = std::pin::pin!(drive_sets(
&mut in_flight,
&mut background,
&mut notifications
));
match select(select(recv, progress), published).await {
Either::Left((Either::Left((Ok(opt), _)), _)) => Step::Message(opt),
Either::Left((Either::Left((Err(e), _)), _)) => break Err(e.into()),
Either::Left((Either::Right((maybe_exec, _)), _)) => {
Step::Progress(maybe_exec.map(Box::new))
}
Either::Right((notification, _)) => Step::Ambient(notification),
}
};
match step {
Step::Progress(Some(exec)) => {
background.push(self.run_task(*exec));
}
Step::Progress(None) => {}
Step::Ambient(notification) => {
if let Err(e) = self
.transport
.send(Message::Notification(notification))
.await
{
tracing::warn!(error = ?e, "failed to send ambient notification");
}
}
Step::Message(Some(Message::Request(request))) => {
if in_flight.len() < max {
in_flight.push(self.handle_request_isolated(request));
} else {
queued.push_back(request);
}
}
Step::Message(Some(Message::Notification(notification))) => {
notifications.push(self.handle_notification(notification));
}
Step::Message(Some(Message::Response(response))) => {
self.state.route_response(response);
}
Step::Message(None) => {
tracing::info!("Connection closed");
break Ok(());
}
}
};
self.state.fail_pending_requests();
while in_flight.next().await.is_some() {}
while background.next().await.is_some() {}
while notifications.next().await.is_some() {}
if let Err(ref err) = outcome {
tracing::error!(error = %err, "Transport error");
}
outcome
}
async fn compute_response(&self, request: &Request) -> Result<serde_json::Value, McpError> {
match request.method.as_ref() {
"initialize" => self.handle_initialize(request).await,
"ping" => self.route_request(request).await,
_ if !self.state.is_initialized() => {
Err(McpError::invalid_request("Server not initialized"))
}
_ => self.route_request(request).await,
}
}
async fn handle_request_isolated(&self, request: Request) -> Option<BackgroundExec> {
use futures::FutureExt;
use std::panic::AssertUnwindSafe;
let id = request.id.clone();
tracing::debug!(method = %request.method, id = %id, "Handling request");
match self.try_begin_task(&request).await {
TaskBegin::Deferred(exec) => return Some(*exec),
TaskBegin::Rejected => return None,
TaskBegin::NotApplicable => {}
}
let computed = AssertUnwindSafe(self.compute_response(&request))
.catch_unwind()
.await;
let response_msg = match computed {
Ok(Ok(result)) => Response::success(id, result),
Ok(Err(e)) => Response::error(id, e.into()),
Err(panic) => {
let detail = panic_message(&*panic);
tracing::error!(method = %request.method, panic = %detail, "Handler panicked");
Response::error(
id,
McpError::internal(format!("handler panicked: {detail}")).into(),
)
}
};
if let Err(e) = self.transport.send(Message::Response(response_msg)).await {
let err: McpError = e.into();
tracing::error!(error = %err, "Failed to send response");
}
None
}
async fn try_begin_task(&self, request: &Request) -> TaskBegin {
if request.method.as_ref() != "tools/call" {
return TaskBegin::NotApplicable;
}
let params = request.params.as_ref();
let Some(task_meta) = params.and_then(|p| p.get("task")) else {
return TaskBegin::NotApplicable;
};
if task_meta.is_null() {
return TaskBegin::NotApplicable;
}
if !self.state.is_initialized() {
return TaskBegin::NotApplicable;
}
let Some(name) = params
.and_then(|p| p.get("name"))
.and_then(|v| v.as_str())
.map(str::to_string)
else {
return TaskBegin::NotApplicable;
};
let args = match params.and_then(|p| p.get("arguments")) {
None => mcpkit_core::types::Object::new(),
Some(serde_json::Value::Object(map)) => map.clone(),
Some(_) => return TaskBegin::NotApplicable,
};
let ttl = task_meta.get("ttl").and_then(serde_json::Value::as_u64);
let client_caps = self.state.client_caps();
let protocol_version = self
.state
.protocol_version()
.unwrap_or(ProtocolVersion::LATEST);
let support = {
let peer = TransportPeer::with_outbound(
self.transport.clone(),
self.state.clone(),
self.config.outbound_request_timeout,
);
let ctx = Context::new(
&request.id,
None,
&client_caps,
&self.state.server_caps,
protocol_version,
&peer,
);
self.server.tool_task_support(&name, &ctx).await
};
if support == mcpkit_core::types::TaskSupport::Forbidden {
let err = McpError::JsonRpc(mcpkit_core::error::JsonRpcError::method_not_found(
format!("tool '{name}' does not support task-augmented execution"),
));
let _ = self
.transport
.send(Message::Response(Response::error(
request.id.clone(),
err.into(),
)))
.await;
return TaskBegin::Rejected;
}
let handle = self.task_store.create(ttl);
let task = handle
.task()
.unwrap_or_else(|| mcpkit_core::types::Task::new(handle.id().clone()));
let create_result =
serde_json::to_value(mcpkit_core::types::CreateTaskResult { task, meta: None })
.unwrap_or_default();
if let Err(e) = self
.transport
.send(Message::Response(Response::success(
request.id.clone(),
create_result,
)))
.await
{
let err: McpError = e.into();
tracing::error!(error = %err, "Failed to send CreateTaskResult");
}
let cancel = handle.cancel_token().unwrap_or_else(CancellationToken::new);
let ctx_data = ContextData::new(
request.id.clone(),
client_caps,
self.state.server_caps.clone(),
protocol_version,
);
TaskBegin::Deferred(Box::new(BackgroundExec {
handle,
name,
args,
ctx_data,
cancel,
}))
}
async fn run_task(&self, exec: BackgroundExec) {
let BackgroundExec {
handle,
name,
args,
ctx_data,
cancel,
} = exec;
let peer = TransportPeer::with_outbound(
self.transport.clone(),
self.state.clone(),
self.config.outbound_request_timeout,
);
let ctx = Context::with_cancellation(
&ctx_data.request_id,
None,
&ctx_data.client_caps,
&ctx_data.server_caps,
ctx_data.protocol_version,
&peer,
cancel,
);
match self.server.call_tool_json(&name, args, &ctx).await {
Ok(payload)
if payload
.get("isError")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false) =>
{
let _ =
handle.fail_with_result(payload, Some("tool reported an error".to_string()));
}
Ok(payload) => {
let _ = handle.complete(payload);
}
Err(e) => {
let _ = handle.fail_with_error(e.into());
}
}
}
async fn route_runtime_tasks(
&self,
method: &str,
params: Option<&serde_json::Value>,
) -> crate::capability::tasks::TaskRoute {
crate::capability::tasks::route_task_store(&self.task_store, method, params).await
}
async fn handle_initialize(&self, request: &Request) -> Result<serde_json::Value, McpError> {
if self.state.is_initialized() {
return Err(McpError::invalid_request("Already initialized"));
}
let params = request
.params
.as_ref()
.ok_or_else(|| McpError::invalid_params("initialize", "missing params"))?;
let requested_version_str = params
.get("protocolVersion")
.and_then(|v| v.as_str())
.unwrap_or("");
let negotiated_version =
ProtocolVersion::negotiate(requested_version_str, ProtocolVersion::ALL)
.unwrap_or(ProtocolVersion::LATEST);
if requested_version_str == negotiated_version.as_str() {
tracing::debug!(
version = %negotiated_version,
"Protocol version negotiated successfully"
);
} else {
tracing::info!(
requested = %requested_version_str,
negotiated = %negotiated_version,
supported = ?ProtocolVersion::ALL.iter().map(ProtocolVersion::as_str).collect::<Vec<_>>(),
"Protocol version negotiation: client requested different version"
);
}
self.state.set_protocol_version(negotiated_version);
if let Some(caps) = params.get("capabilities") {
if let Ok(client_caps) = serde_json::from_value::<ClientCapabilities>(caps.clone()) {
self.state.set_client_caps(client_caps);
}
}
let result = serde_json::json!({
"protocolVersion": negotiated_version.as_str(),
"serverInfo": self.server.server_info(),
"capabilities": self.state.server_caps
});
self.state.set_initialized();
Ok(result)
}
async fn route_request(&self, request: &Request) -> Result<serde_json::Value, McpError> {
let method = request.method.as_ref();
let params = request.params.as_ref();
let unowned_task: Option<McpError> = match self.route_runtime_tasks(method, params).await {
crate::capability::tasks::TaskRoute::Handled(result) => return result,
crate::capability::tasks::TaskRoute::NotTaskMethod => None,
unowned => unowned.or_unknown_task().and_then(Result::err),
};
let progress_token = extract_progress_token(params);
let peer = TransportPeer::with_outbound(
self.transport.clone(),
self.state.clone(),
self.config.outbound_request_timeout,
);
let client_caps = self.state.client_caps();
let protocol_version = self
.state
.protocol_version()
.unwrap_or(ProtocolVersion::LATEST);
let cancel = CancellationToken::new();
let cancel_key = request.id.to_string();
self.state
.register_cancellation(&cancel_key, cancel.clone());
let ctx = Context::with_cancellation(
&request.id,
progress_token.as_ref(),
&client_caps,
&self.state.server_caps,
protocol_version,
&peer,
cancel,
);
let result = self.server.route(method, params, &ctx).await;
self.state.remove_cancellation(&cancel_key);
if let Some(unowned) = unowned_task {
if result.as_ref().err().map(McpError::code)
== Some(mcpkit_core::error::codes::METHOD_NOT_FOUND)
{
return Err(unowned);
}
}
result
}
async fn handle_notification(&self, notification: Notification) -> Result<(), McpError> {
let method = notification.method.as_ref();
tracing::debug!(method = %method, "Handling notification");
if method == crate::router::notifications::CANCELLED {
if let Some(request_id) = notification
.params
.as_ref()
.and_then(|p| {
serde_json::from_value::<mcpkit_core::types::CancelledNotificationParams>(
p.clone(),
)
.ok()
})
.and_then(|c| c.request_id)
{
self.state.cancel_request(&request_id.to_string());
}
return Ok(());
}
let client_caps = self.state.client_caps();
let protocol_version = self
.state
.protocol_version()
.unwrap_or(ProtocolVersion::LATEST);
let peer = TransportPeer::with_outbound(
self.transport.clone(),
self.state.clone(),
self.config.outbound_request_timeout,
);
let ctx = Context::for_notification(
&client_caps,
&self.state.server_caps,
protocol_version,
&peer,
);
self.server
.route_notification(method, notification.params.as_ref(), &ctx)
.await;
Ok(())
}
}
impl<H, T, R, P, K, Tr> ServerRuntime<Server<H, T, R, P, K>, Tr>
where
H: ServerHandler + Send + Sync,
T: Send + Sync,
R: Send + Sync,
P: Send + Sync,
K: Send + Sync,
Tr: Transport + 'static,
Tr::Error: Into<McpError>,
{
pub fn new(server: Server<H, T, R, P, K>, transport: Tr) -> Self {
Self::with_config(server, transport, RuntimeConfig::default())
}
pub fn with_config(
server: Server<H, T, R, P, K>,
transport: Tr,
config: RuntimeConfig,
) -> Self {
let caps = server.capabilities().clone();
let task_store = Arc::new(
crate::capability::tasks::TaskManager::with_default_ttl(config.default_task_ttl_ms)
.with_poll_interval(config.default_task_poll_interval_ms),
);
let state = Arc::new(ServerState::new(caps));
if config.task_status_notifications {
let _ = task_store.set_observer(Arc::new(
crate::capability::tasks::TaskStatusNotifier::new(state.clone()),
));
}
Self {
server,
transport: Arc::new(transport),
state,
task_store,
config,
}
}
}
#[allow(async_fn_in_trait)]
pub trait RequestRouter: Send + Sync {
fn server_info(&self) -> mcpkit_core::capability::ServerInfo;
async fn route(
&self,
method: &str,
params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError>;
async fn route_notification(
&self,
_method: &str,
_params: Option<&serde_json::Value>,
_ctx: &Context<'_>,
) {
}
async fn tool_task_support(
&self,
_name: &str,
_ctx: &Context<'_>,
) -> mcpkit_core::types::TaskSupport {
mcpkit_core::types::TaskSupport::Forbidden
}
async fn call_tool_json(
&self,
name: &str,
_args: mcpkit_core::types::Object,
_ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
Err(McpError::method_not_found(name))
}
}
impl<H, T, R, P, K> Server<H, T, R, P, K>
where
H: ServerHandler + Send + Sync + 'static,
T: Send + Sync + 'static,
R: Send + Sync + 'static,
P: Send + Sync + 'static,
K: Send + Sync + 'static,
Self: RequestRouter,
{
pub async fn serve<Tr>(self, transport: Tr) -> Result<(), McpError>
where
Tr: Transport + 'static,
Tr::Error: Into<McpError>,
{
let runtime = ServerRuntime::new(self, transport);
runtime.run().await
}
}
impl<H, T, R, P, K> RequestRouter for Server<H, T, R, P, K>
where
H: ServerHandler + Send + Sync,
T: ToolSlot,
R: ResourceSlot,
P: PromptSlot,
K: TaskSlot,
{
fn server_info(&self) -> mcpkit_core::capability::ServerInfo {
self.handler().server_info()
}
async fn route_notification(
&self,
method: &str,
_params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) {
crate::router::dispatch_notification_hooks(self.handler(), method, ctx).await;
}
async fn route(
&self,
method: &str,
params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
if method == "ping" {
return Ok(serde_json::json!({}));
}
let page_size = self.list_page_size;
if let Some(handler) = self.tools.as_tool_handler() {
if let Some(result) = route_tools(handler, method, params, ctx, page_size).await {
return result;
}
}
if let Some(handler) = self.resources.as_resource_handler() {
if let Some(result) = route_resources(handler, method, params, ctx, page_size).await {
return result;
}
}
if let Some(handler) = self.prompts.as_prompt_handler() {
if let Some(result) = route_prompts(handler, method, params, ctx, page_size).await {
return result;
}
}
if let Some(handler) = self.tasks.as_task_handler() {
if let Some(result) = route_tasks(handler, method, params, ctx).await {
return result;
}
}
if let Some(result) =
crate::router::route_logging(self.handler(), self.capabilities(), method, params, ctx)
.await
{
return result;
}
if let Some(result) =
crate::router::route_completion(self.completion.as_deref(), method, params, ctx).await
{
return result;
}
Err(McpError::method_not_found(method))
}
async fn tool_task_support(
&self,
name: &str,
ctx: &Context<'_>,
) -> mcpkit_core::types::TaskSupport {
match self.tools.as_tool_handler() {
Some(handler) => crate::router::tool_task_support(handler, name, ctx).await,
None => mcpkit_core::types::TaskSupport::Forbidden,
}
}
async fn call_tool_json(
&self,
name: &str,
args: mcpkit_core::types::Object,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match self.tools.as_tool_handler() {
Some(handler) => crate::router::call_tool_json(handler, name, args, ctx).await,
None => Err(McpError::method_not_found(name)),
}
}
}
fn panic_message(panic: &(dyn std::any::Any + Send)) -> String {
if let Some(s) = panic.downcast_ref::<&str>() {
(*s).to_string()
} else if let Some(s) = panic.downcast_ref::<String>() {
s.clone()
} else {
"unknown panic".to_string()
}
}
fn extract_progress_token(params: Option<&serde_json::Value>) -> Option<ProgressToken> {
params.and_then(mcpkit_core::types::Meta::progress_token_from_params)
}
#[cfg(test)]
mod tests {
use super::*;
use mcpkit_core::capability::{ClientCapabilities, ServerInfo};
use mcpkit_core::protocol::RequestId;
use mcpkit_core::types::content::Role;
use mcpkit_core::types::elicitation::ElicitRequest;
use mcpkit_core::types::sampling::{CreateMessageRequest, CreateMessageResult};
use mcpkit_transport::MemoryTransport;
use std::time::Duration;
use tokio::sync::Notify;
use tokio::time::timeout;
struct PanicRouter;
impl RequestRouter for PanicRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("panic-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"panic" => panic!("boom in handler"),
"ok" => Ok(serde_json::json!("ok")),
other => Err(McpError::method_not_found(other)),
}
}
}
struct CoordRouter {
started: Arc<Notify>,
release: Arc<Notify>,
}
impl RequestRouter for CoordRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("coord-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"blocker" => {
self.started.notify_one();
self.release.notified().await;
Ok(serde_json::json!("blocked-done"))
}
"fast" => Ok(serde_json::json!("fast-done")),
other => Err(McpError::method_not_found(other)),
}
}
}
struct PingRouter;
impl RequestRouter for PingRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("ping-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"ping" => Ok(serde_json::json!({})),
other => Err(McpError::method_not_found(other)),
}
}
}
struct CancelRouter {
started: Arc<Notify>,
}
impl RequestRouter for CancelRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("cancel-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"wait_cancel" => {
self.started.notify_one();
ctx.cancelled().await;
Ok(serde_json::json!(ctx.is_cancelled()))
}
other => Err(McpError::method_not_found(other)),
}
}
}
struct OutboundRouter;
impl RequestRouter for OutboundRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("outbound-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"ask" => ctx.request("ask/upstream", None).await,
other => Err(McpError::method_not_found(other)),
}
}
}
struct ElicitRouter;
impl RequestRouter for ElicitRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("elicit-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"ask_name" => {
let result = ctx
.elicit(ElicitRequest::text("Your name?", "name"))
.await?;
Ok(serde_json::json!({
"accepted": result.is_accepted(),
"name": result.get_string("name"),
}))
}
other => Err(McpError::method_not_found(other)),
}
}
}
struct SampleRouter;
impl RequestRouter for SampleRouter {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("sample-test", "0.0.0")
}
async fn route(
&self,
method: &str,
_params: Option<&serde_json::Value>,
ctx: &Context<'_>,
) -> Result<serde_json::Value, McpError> {
match method {
"summarize" => {
let result = ctx
.create_message(CreateMessageRequest::simple("hello", 100))
.await?;
Ok(serde_json::json!({ "text": result.as_text() }))
}
other => Err(McpError::method_not_found(other)),
}
}
}
fn req(method: &'static str, id: u64) -> Message {
Message::Request(Request::new(method, id))
}
async fn next_response(transport: &MemoryTransport) -> Response {
for _ in 0..16 {
let msg = timeout(Duration::from_secs(2), transport.recv())
.await
.expect("no response (connection died?)")
.expect("recv ok")
.expect("some message");
match msg {
Message::Response(r) => return r,
Message::Notification(_) => continue,
other => panic!("expected response, got {other:?}"),
}
}
panic!("no response after 16 messages");
}
fn notif_msg(method: &str) -> Message {
Message::Notification(Notification::with_params(
method.to_string(),
serde_json::json!({}),
))
}
struct RootsHookHandler {
initialized: Arc<std::sync::atomic::AtomicBool>,
roots_changed: Arc<std::sync::atomic::AtomicUsize>,
seen_roots: Arc<std::sync::Mutex<Vec<mcpkit_core::types::Root>>>,
done: Arc<Notify>,
}
impl crate::handler::ServerHandler for RootsHookHandler {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("roots-test", "0.0.0")
}
async fn on_initialized(&self, _ctx: &Context<'_>) {
self.initialized
.store(true, std::sync::atomic::Ordering::SeqCst);
self.done.notify_one();
}
async fn on_roots_list_changed(&self, ctx: &Context<'_>) {
self.roots_changed
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
if let Ok(roots) = ctx.list_roots().await {
*self.seen_roots.lock().expect("lock") = roots;
}
self.done.notify_one();
}
}
#[tokio::test]
async fn notification_hooks_fire_and_on_roots_list_changed_can_list_roots() {
use crate::builder::ServerBuilder;
use mcpkit_core::capability::ClientCapabilities;
use mcpkit_core::types::{ListRootsResult, Root};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
let initialized = Arc::new(AtomicBool::new(false));
let roots_changed = Arc::new(AtomicUsize::new(0));
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let done = Arc::new(Notify::new());
let handler = RootsHookHandler {
initialized: initialized.clone(),
roots_changed: roots_changed.clone(),
seen_roots: seen.clone(),
done: done.clone(),
};
let (client, server_tr) = MemoryTransport::pair();
let runtime = ServerRuntime::new(ServerBuilder::new(handler).build(), server_tr);
runtime.state().set_initialized();
runtime
.state()
.set_client_caps(ClientCapabilities::default().with_roots());
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(notif_msg("notifications/initialized"))
.await
.expect("send");
timeout(Duration::from_secs(2), done.notified())
.await
.expect("on_initialized never ran");
assert!(initialized.load(Ordering::SeqCst));
client
.send(notif_msg("notifications/roots/list_changed"))
.await
.expect("send");
let roots_req = match timeout(Duration::from_secs(2), client.recv())
.await
.expect("no roots/list request")
.expect("recv ok")
.expect("some message")
{
Message::Request(r) => r,
other => panic!("expected roots/list, got {other:?}"),
};
assert_eq!(roots_req.method.as_ref(), "roots/list");
let result = ListRootsResult {
roots: vec![Root::new("file:///work")],
meta: None,
};
client
.send(Message::Response(Response::success(
roots_req.id.clone(),
serde_json::to_value(result).expect("serialize"),
)))
.await
.expect("send");
timeout(Duration::from_secs(2), done.notified())
.await
.expect("on_roots_list_changed never finished");
assert_eq!(roots_changed.load(Ordering::SeqCst), 1);
assert_eq!(seen.lock().expect("lock")[0].uri, "file:///work");
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn roots_list_changed_is_ignored_without_roots_capability() {
use crate::builder::ServerBuilder;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
let roots_changed = Arc::new(AtomicUsize::new(0));
let handler = RootsHookHandler {
initialized: Arc::new(AtomicBool::new(false)),
roots_changed: roots_changed.clone(),
seen_roots: Arc::new(std::sync::Mutex::new(Vec::new())),
done: Arc::new(Notify::new()),
};
let (client, server_tr) = MemoryTransport::pair();
let runtime = ServerRuntime::new(ServerBuilder::new(handler).build(), server_tr);
runtime.state().set_initialized();
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(notif_msg("notifications/roots/list_changed"))
.await
.expect("send");
client.send(req("ping", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert_eq!(
roots_changed.load(Ordering::SeqCst),
0,
"on_roots_list_changed must not fire without the roots capability"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn task_transition_publishes_status_notification() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let task_store = Arc::new(crate::capability::tasks::TaskManager::new());
task_store
.set_observer(Arc::new(crate::capability::tasks::TaskStatusNotifier::new(
state.clone(),
)))
.expect("install observer");
let runtime = ServerRuntime {
server: PingRouter,
transport: Arc::new(server),
state,
task_store: Arc::clone(&task_store),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
let task = task_store.create(None);
task.complete(serde_json::json!({"ok": true}))
.expect("complete");
let msg = timeout(Duration::from_secs(2), client.recv())
.await
.expect("no notification (loop never drained the queue?)")
.expect("recv ok")
.expect("some message");
let Message::Notification(notification) = msg else {
panic!("expected notification, got {msg:?}");
};
assert_eq!(notification.method, "notifications/tasks/status");
let params = notification.params.expect("params");
assert_eq!(params["taskId"], task.id().as_str());
assert_eq!(params["status"], "completed");
assert!(
params.get("_meta").is_none(),
"status notification must not carry _meta: {params}"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn task_status_notifications_can_be_disabled() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let config = RuntimeConfig::new().task_status_notifications(false);
let task_store = Arc::new(crate::capability::tasks::TaskManager::new());
let runtime = ServerRuntime {
server: PingRouter,
transport: Arc::new(server),
state,
task_store: Arc::clone(&task_store),
config,
};
let handle = tokio::spawn(async move { runtime.run().await });
let task = task_store.create(None);
task.complete(serde_json::json!({})).expect("complete");
client.send(req("ping", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn panic_in_handler_returns_internal_error_and_keeps_connection() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: PanicRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("panic", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
let err = resp.error.expect("expected error response");
assert!(
err.message.contains("panicked"),
"unexpected error message: {}",
err.message
);
client.send(req("ok", 2)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(2));
assert!(
resp.result.is_some(),
"expected success after a prior panic"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn ping_is_answered_before_initialize() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
let runtime = ServerRuntime {
server: PingRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("ping", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert!(
resp.error.is_none(),
"ping before initialize must not error: {:?}",
resp.error
);
assert!(resp.result.is_some(), "ping should return a result");
client.send(req("tools/list", 2)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(2));
assert!(
resp.error.is_some(),
"non-ping requests before initialize must still be rejected"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn requests_are_processed_concurrently() {
let (client, server) = MemoryTransport::pair();
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: CoordRouter {
started: started.clone(),
release: release.clone(),
},
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("blocker", 1)).await.expect("send");
client.send(req("fast", 2)).await.expect("send");
timeout(Duration::from_secs(2), started.notified())
.await
.expect("blocker never started");
let resp = next_response(&client).await;
assert_eq!(
resp.id,
RequestId::Number(2),
"fast request should finish first"
);
release.notify_one();
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn max_concurrent_requests_limits_in_flight() {
let (client, server) = MemoryTransport::pair();
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: CoordRouter {
started: started.clone(),
release: release.clone(),
},
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig {
auto_initialized: true,
max_concurrent_requests: 1,
..RuntimeConfig::default()
},
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("blocker", 1)).await.expect("send");
client.send(req("fast", 2)).await.expect("send");
timeout(Duration::from_secs(2), started.notified())
.await
.expect("blocker never started");
let early = timeout(Duration::from_millis(200), client.recv()).await;
assert!(
early.is_err(),
"fast request was processed despite max_concurrent_requests = 1"
);
release.notify_one();
assert_eq!(next_response(&client).await.id, RequestId::Number(1));
assert_eq!(next_response(&client).await.id, RequestId::Number(2));
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn cancelled_notification_trips_in_flight_handler() {
let (client, server) = MemoryTransport::pair();
let started = Arc::new(Notify::new());
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: CancelRouter {
started: started.clone(),
},
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("wait_cancel", 1)).await.expect("send");
timeout(Duration::from_secs(2), started.notified())
.await
.expect("handler never started");
let cancel = Message::Notification(Notification::with_params(
"notifications/cancelled".to_string(),
serde_json::json!({ "requestId": 1 }),
));
client.send(cancel).await.expect("send cancel");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert_eq!(
resp.result,
Some(serde_json::json!(true)),
"ctx.is_cancelled() should be true after notifications/cancelled"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn notifier_sends_list_changed_outside_request() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
let runtime = ServerRuntime {
server: PingRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let notifier = runtime.notifier();
notifier.tools_list_changed().await.expect("notify");
let msg = timeout(Duration::from_secs(2), client.recv())
.await
.expect("no notification (timed out)")
.expect("recv ok")
.expect("some message");
match msg {
Message::Notification(n) => {
assert_eq!(n.method.as_ref(), "notifications/tools/list_changed");
}
other => panic!("expected a notification, got {other:?}"),
}
}
#[tokio::test]
async fn server_initiated_request_roundtrips_at_concurrency_limit() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: OutboundRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig {
auto_initialized: true,
max_concurrent_requests: 1,
..RuntimeConfig::default()
},
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("ask", 1)).await.expect("send");
let outbound = match timeout(Duration::from_secs(2), client.recv())
.await
.expect("no outbound request (timed out)")
.expect("recv ok")
.expect("some message")
{
Message::Request(r) => r,
other => panic!("expected a server-initiated request, got {other:?}"),
};
assert_eq!(outbound.method.as_ref(), "ask/upstream");
client
.send(Message::Response(Response::success(
outbound.id.clone(),
serde_json::json!({ "answer": 42 }),
)))
.await
.expect("send response");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert_eq!(resp.result, Some(serde_json::json!({ "answer": 42 })));
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn server_initiated_request_times_out() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: OutboundRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig {
outbound_request_timeout: Duration::from_millis(100),
..RuntimeConfig::default()
},
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("ask", 1)).await.expect("send");
let _outbound = timeout(Duration::from_secs(2), client.recv())
.await
.expect("no outbound request")
.expect("recv ok")
.expect("some message");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert!(resp.error.is_some(), "timed-out request should error");
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn ctx_elicit_roundtrips() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
state.set_client_caps(ClientCapabilities::default().with_elicitation());
let runtime = ServerRuntime {
server: ElicitRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("ask_name", 1)).await.expect("send");
let elicit = match timeout(Duration::from_secs(2), client.recv())
.await
.expect("no elicitation request")
.expect("recv ok")
.expect("some message")
{
Message::Request(r) => r,
other => panic!("expected elicitation/create, got {other:?}"),
};
assert_eq!(elicit.method.as_ref(), "elicitation/create");
assert!(
elicit
.params
.as_ref()
.and_then(|p| p.get("requestedSchema"))
.is_some(),
"elicitation request should carry a requestedSchema"
);
client
.send(Message::Response(Response::success(
elicit.id.clone(),
serde_json::json!({ "action": "accept", "content": { "name": "Ada" } }),
)))
.await
.expect("send response");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert_eq!(
resp.result,
Some(serde_json::json!({ "accepted": true, "name": "Ada" }))
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn ctx_elicit_requires_client_capability() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: ElicitRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("ask_name", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert!(
resp.error.is_some(),
"elicit without client capability should error"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn ctx_create_message_roundtrips() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
state.set_client_caps(ClientCapabilities::default().with_sampling());
let runtime = ServerRuntime {
server: SampleRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("summarize", 1)).await.expect("send");
let sampling = match timeout(Duration::from_secs(2), client.recv())
.await
.expect("no sampling request")
.expect("recv ok")
.expect("some message")
{
Message::Request(r) => r,
other => panic!("expected sampling/createMessage, got {other:?}"),
};
assert_eq!(sampling.method.as_ref(), "sampling/createMessage");
let result = CreateMessageResult {
role: Role::Assistant,
content: mcpkit_core::types::OneOrMany::One(mcpkit_core::types::SamplingContent::text(
"a summary",
)),
model: "test-model".to_string(),
stop_reason: None,
meta: None,
};
client
.send(Message::Response(Response::success(
sampling.id.clone(),
serde_json::to_value(result).expect("serialize result"),
)))
.await
.expect("send response");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert_eq!(
resp.result,
Some(serde_json::json!({ "text": "a summary" }))
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn ctx_create_message_requires_client_capability() {
let (client, server) = MemoryTransport::pair();
let state = Arc::new(ServerState::new(ServerCapabilities::default()));
state.set_initialized();
let runtime = ServerRuntime {
server: SampleRouter,
transport: Arc::new(server),
state,
task_store: Arc::new(crate::capability::tasks::TaskManager::new()),
config: RuntimeConfig::default(),
};
let handle = tokio::spawn(async move { runtime.run().await });
client.send(req("summarize", 1)).await.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert!(
resp.error.is_some(),
"create_message without client capability should error"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[test]
fn test_server_state_initialization() {
let state = ServerState::new(ServerCapabilities::default());
assert!(!state.is_initialized());
state.set_initialized();
assert!(state.is_initialized());
}
#[test]
fn test_cancellation_management() {
let state = ServerState::new(ServerCapabilities::default());
let token = CancellationToken::new();
state.register_cancellation("req-1", token.clone());
assert!(!token.is_cancelled());
state.cancel_request("req-1");
assert!(token.is_cancelled());
state.remove_cancellation("req-1");
}
#[test]
fn test_runtime_config_default() {
let config = RuntimeConfig::default();
assert!(config.auto_initialized);
assert_eq!(config.max_concurrent_requests, 100);
}
#[test]
fn test_extract_progress_token_string() -> Result<(), Box<dyn std::error::Error>> {
let params = serde_json::json!({
"_meta": {
"progressToken": "my-token-123"
},
"name": "test-tool"
});
let token = extract_progress_token(Some(¶ms));
assert!(token.is_some());
assert_eq!(
token.ok_or("Token not found")?,
ProgressToken::String("my-token-123".to_string())
);
Ok(())
}
#[test]
fn test_extract_progress_token_number() -> Result<(), Box<dyn std::error::Error>> {
let params = serde_json::json!({
"_meta": {
"progressToken": 42
},
"arguments": {}
});
let token = extract_progress_token(Some(¶ms));
assert!(token.is_some());
assert_eq!(token.ok_or("Token not found")?, ProgressToken::Number(42));
Ok(())
}
#[test]
fn test_extract_progress_token_missing_meta() {
let params = serde_json::json!({
"name": "test-tool",
"arguments": {}
});
let token = extract_progress_token(Some(¶ms));
assert!(token.is_none());
}
#[test]
fn test_extract_progress_token_missing_token() {
let params = serde_json::json!({
"_meta": {},
"name": "test-tool"
});
let token = extract_progress_token(Some(¶ms));
assert!(token.is_none());
}
#[test]
fn test_extract_progress_token_none_params() {
let token = extract_progress_token(None);
assert!(token.is_none());
}
#[tokio::test]
async fn task_augmented_tools_call_runs_in_background() {
use crate::builder::ServerBuilder;
use crate::handler::{ServerHandler, ToolHandler};
use mcpkit_core::protocol::Request;
use mcpkit_core::types::{TaskSupport, Tool, ToolOutput};
struct H;
impl ServerHandler for H {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("t", "1.0.0")
}
}
impl ToolHandler for H {
async fn list_tools(&self, _ctx: &Context<'_>) -> Result<Vec<Tool>, McpError> {
Ok(vec![Tool::new("slow").task_support(TaskSupport::Optional)])
}
async fn call_tool(
&self,
name: &str,
_args: serde_json::Map<String, serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<ToolOutput, McpError> {
Ok(ToolOutput::text(format!("done:{name}")))
}
}
let request = |id: u64, method: &'static str, params: serde_json::Value| {
Message::Request(Request {
jsonrpc: "2.0".into(),
id: RequestId::Number(id),
method: method.into(),
params: Some(params),
})
};
let (client, server_tr) = MemoryTransport::pair();
let built = ServerBuilder::new(H).with_tools(H).build();
let runtime = ServerRuntime::new(built, server_tr);
runtime.state().set_initialized();
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(request(
1,
"tools/call",
serde_json::json!({ "name": "slow", "arguments": {}, "task": {} }),
))
.await
.expect("send");
let resp = next_response(&client).await;
assert_eq!(resp.id, RequestId::Number(1));
assert!(
resp.error.is_none(),
"augmented call errored: {:?}",
resp.error
);
let result = resp.result.expect("create result");
assert_eq!(result["task"]["status"], "working");
let task_id = result["task"]["taskId"]
.as_str()
.expect("taskId")
.to_string();
let mut payload = None;
for attempt in 0..100u64 {
client
.send(request(
100 + attempt,
"tasks/result",
serde_json::json!({ "taskId": task_id }),
))
.await
.expect("send");
let r = next_response(&client).await;
if r.error.is_none() {
payload = r.result;
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
let payload = payload.expect("task completed with a payload");
assert!(
payload["content"][0]["text"]
.as_str()
.unwrap_or_default()
.contains("done:slow"),
"unexpected task payload: {payload}"
);
client
.send(request(
999,
"tasks/get",
serde_json::json!({ "taskId": task_id }),
))
.await
.expect("send");
let got = next_response(&client).await;
assert_eq!(got.result.expect("task")["status"], "completed");
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn tasks_result_blocks_while_loop_stays_live() {
use crate::builder::ServerBuilder;
use crate::handler::{ServerHandler, ToolHandler};
use mcpkit_core::protocol::Request;
use mcpkit_core::types::{TaskSupport, Tool, ToolOutput};
struct H(Arc<tokio::sync::Notify>);
impl ServerHandler for H {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("t", "1.0.0")
}
}
impl ToolHandler for H {
async fn list_tools(&self, _ctx: &Context<'_>) -> Result<Vec<Tool>, McpError> {
Ok(vec![Tool::new("gated").task_support(TaskSupport::Optional)])
}
async fn call_tool(
&self,
_name: &str,
_args: serde_json::Map<String, serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<ToolOutput, McpError> {
self.0.notified().await;
Ok(ToolOutput::text("released"))
}
}
let request = |id: u64, method: &'static str, params: serde_json::Value| {
Message::Request(Request {
jsonrpc: "2.0".into(),
id: RequestId::Number(id),
method: method.into(),
params: Some(params),
})
};
let release = Arc::new(tokio::sync::Notify::new());
let (client, server_tr) = MemoryTransport::pair();
let built = ServerBuilder::new(H(release.clone()))
.with_tools(H(release.clone()))
.build();
let runtime = ServerRuntime::new(built, server_tr);
runtime.state().set_initialized();
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(request(
1,
"tools/call",
serde_json::json!({ "name": "gated", "arguments": {}, "task": {} }),
))
.await
.expect("send");
let resp = next_response(&client).await;
let task_id = resp.result.expect("create result")["task"]["taskId"]
.as_str()
.expect("taskId")
.to_string();
client
.send(request(
2,
"tasks/result",
serde_json::json!({ "taskId": task_id }),
))
.await
.expect("send");
client
.send(request(3, "tools/list", serde_json::json!({})))
.await
.expect("send");
let live = timeout(Duration::from_secs(2), next_response(&client))
.await
.expect("loop stalled while tasks/result was blocking");
assert_eq!(live.id, RequestId::Number(3), "expected tools/list reply");
release.notify_one();
let result = timeout(Duration::from_secs(2), next_response(&client))
.await
.expect("blocked tasks/result never completed");
assert_eq!(result.id, RequestId::Number(2));
let payload = result.result.expect("payload");
assert!(
payload["content"][0]["text"]
.as_str()
.unwrap_or_default()
.contains("released"),
"unexpected payload: {payload}"
);
assert_eq!(
payload["_meta"]["io.modelcontextprotocol/related-task"]["taskId"],
task_id.as_str()
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn task_augmented_call_on_forbidden_tool_is_rejected() {
use crate::builder::ServerBuilder;
use crate::handler::{ServerHandler, ToolHandler};
use mcpkit_core::protocol::Request;
use mcpkit_core::types::{Tool, ToolOutput};
struct H;
impl ServerHandler for H {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("t", "1.0.0")
}
}
impl ToolHandler for H {
async fn list_tools(&self, _ctx: &Context<'_>) -> Result<Vec<Tool>, McpError> {
Ok(vec![Tool::new("plain")])
}
async fn call_tool(
&self,
_name: &str,
_args: serde_json::Map<String, serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<ToolOutput, McpError> {
Ok(ToolOutput::text("ok"))
}
}
let (client, server_tr) = MemoryTransport::pair();
let runtime = ServerRuntime::new(ServerBuilder::new(H).with_tools(H).build(), server_tr);
runtime.state().set_initialized();
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(Message::Request(Request {
jsonrpc: "2.0".into(),
id: RequestId::Number(1),
method: "tools/call".into(),
params: Some(serde_json::json!({ "name": "plain", "task": {} })),
}))
.await
.expect("send");
let resp = next_response(&client).await;
let err = resp
.error
.expect("a forbidden tool must reject task augmentation");
assert_eq!(err.code, -32601, "wrong rejection code: {err:?}");
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[cfg(feature = "schema-validation")]
#[tokio::test]
async fn task_path_validates_input_via_decorator() {
use crate::builder::ServerBuilder;
use crate::handler::{ServerHandler, ToolHandler};
use mcpkit_core::protocol::Request;
use mcpkit_core::types::{TaskSupport, Tool, ToolOutput};
struct H;
impl ServerHandler for H {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("t", "1.0.0")
}
}
impl ToolHandler for H {
async fn list_tools(&self, _ctx: &Context<'_>) -> Result<Vec<Tool>, McpError> {
Ok(vec![
Tool::new("slow")
.task_support(TaskSupport::Optional)
.input_schema(serde_json::json!({
"type": "object",
"properties": { "n": { "type": "number" } },
"required": ["n"]
})),
])
}
async fn call_tool(
&self,
name: &str,
_args: serde_json::Map<String, serde_json::Value>,
_ctx: &Context<'_>,
) -> Result<ToolOutput, McpError> {
Ok(ToolOutput::text(format!("done:{name}")))
}
}
let request = |id: u64, method: &'static str, params: serde_json::Value| {
Message::Request(Request {
jsonrpc: "2.0".into(),
id: RequestId::Number(id),
method: method.into(),
params: Some(params),
})
};
let (client, server_tr) = MemoryTransport::pair();
let built = ServerBuilder::new(H)
.with_tools(H)
.validate_tool_io()
.build();
let runtime = ServerRuntime::new(built, server_tr);
runtime.state().set_initialized();
let handle = tokio::spawn(async move { runtime.run().await });
client
.send(request(
1,
"tools/call",
serde_json::json!({ "name": "slow", "arguments": {}, "task": {} }),
))
.await
.expect("send");
let resp = next_response(&client).await;
let task_id = resp.result.expect("create result")["task"]["taskId"]
.as_str()
.expect("taskId")
.to_string();
let mut payload = None;
for attempt in 0..100u64 {
client
.send(request(
100 + attempt,
"tasks/result",
serde_json::json!({ "taskId": task_id }),
))
.await
.expect("send");
let r = next_response(&client).await;
if r.error.is_none() {
payload = r.result;
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
let payload = payload.expect("task completed with a payload");
assert_eq!(
payload["isError"],
serde_json::json!(true),
"task path must validate input: {payload}"
);
assert!(
!payload["content"][0]["text"]
.as_str()
.unwrap_or_default()
.contains("done:slow"),
"the tool body must not have run: {payload}"
);
drop(client);
let _ = timeout(Duration::from_secs(2), handle).await;
}
#[tokio::test]
async fn logging_set_level_dispatches_when_advertised_else_method_not_found() {
use crate::builder::ServerBuilder;
use crate::context::NoOpPeer;
use crate::handler::ServerHandler;
use mcpkit_core::capability::{ClientCapabilities, ServerCapabilities};
use mcpkit_core::protocol::RequestId;
use mcpkit_core::protocol_version::ProtocolVersion;
use mcpkit_core::types::LoggingLevel;
use std::sync::Mutex;
struct H(Arc<Mutex<Option<LoggingLevel>>>);
impl ServerHandler for H {
fn server_info(&self) -> ServerInfo {
ServerInfo::new("t", "1.0.0")
}
async fn set_log_level(
&self,
level: LoggingLevel,
_ctx: &Context<'_>,
) -> Result<(), McpError> {
*self.0.lock().unwrap() = Some(level);
Ok(())
}
}
let request_id = RequestId::Number(1);
let client_caps = ClientCapabilities::default();
let server_caps = ServerCapabilities::default();
let peer = NoOpPeer;
let ctx = Context::new(
&request_id,
None,
&client_caps,
&server_caps,
ProtocolVersion::LATEST,
&peer,
);
let seen = Arc::new(Mutex::new(None));
let server = ServerBuilder::new(H(Arc::clone(&seen)))
.capabilities(ServerCapabilities::new().with_logging())
.build();
let out = server
.route(
"logging/setLevel",
Some(&serde_json::json!({ "level": "warning" })),
&ctx,
)
.await
.expect("setLevel dispatched");
assert_eq!(out, serde_json::json!({}));
assert_eq!(*seen.lock().unwrap(), Some(LoggingLevel::Warning));
assert!(
server
.route(
"logging/setLevel",
Some(&serde_json::json!({ "level": "loud" })),
&ctx,
)
.await
.is_err()
);
let plain = ServerBuilder::new(H(Arc::new(Mutex::new(None)))).build();
let err = plain
.route(
"logging/setLevel",
Some(&serde_json::json!({ "level": "info" })),
&ctx,
)
.await
.expect_err("no logging capability -> method not found");
assert!(matches!(err, McpError::MethodNotFound { .. }));
}
#[tokio::test]
async fn context_log_emits_message_notification() {
use crate::context::Peer;
use mcpkit_core::capability::{ClientCapabilities, ServerCapabilities};
use mcpkit_core::protocol::RequestId;
use mcpkit_core::protocol_version::ProtocolVersion;
use mcpkit_core::types::LoggingLevel;
use std::pin::Pin;
use std::sync::Mutex;
struct RecPeer(Arc<Mutex<Vec<Notification>>>);
impl Peer for RecPeer {
fn notify(
&self,
notification: Notification,
) -> Pin<Box<dyn std::future::Future<Output = Result<(), McpError>> + Send + '_>>
{
self.0.lock().unwrap().push(notification);
Box::pin(async { Ok(()) })
}
}
let seen = Arc::new(Mutex::new(Vec::new()));
let peer = RecPeer(Arc::clone(&seen));
let request_id = RequestId::Number(1);
let client_caps = ClientCapabilities::default();
let server_caps = ServerCapabilities::default();
let ctx = Context::new(
&request_id,
None,
&client_caps,
&server_caps,
ProtocolVersion::LATEST,
&peer,
);
ctx.log(LoggingLevel::Error, Some("db"), serde_json::json!("boom"))
.await
.expect("log sent");
let seen = seen.lock().unwrap();
assert_eq!(seen.len(), 1);
assert_eq!(seen[0].method.as_ref(), "notifications/message");
let params = seen[0].params.as_ref().expect("params");
assert_eq!(params["level"], serde_json::json!("error"));
assert_eq!(params["logger"], serde_json::json!("db"));
assert_eq!(params["data"], serde_json::json!("boom"));
}
}