use std::sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
};
use futures::{
channel::oneshot,
future::{BoxFuture, FutureExt, Shared},
};
use serde_json::{Map, Value};
use super::McpConnectionTo;
use crate::{
Error, RequestCancellation, Role,
schema::v1::{McpError, McpRequestId, McpServerAcpId},
};
#[derive(Debug)]
pub struct McpRequest {
pub method: String,
pub params: Option<Map<String, Value>>,
}
#[derive(Debug)]
pub enum McpOutcome {
Result(Value),
Error(McpError),
}
type Notify = dyn Fn(String, Option<Map<String, Value>>) -> BoxFuture<'static, Result<(), Error>>
+ Send
+ Sync;
#[derive(Clone)]
pub struct McpOperationCancellation {
state: Arc<CancellationState>,
}
struct CancellationState {
cancelled: AtomicBool,
sender: Mutex<Option<oneshot::Sender<()>>>,
signal: Shared<BoxFuture<'static, ()>>,
}
impl std::fmt::Debug for McpOperationCancellation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpOperationCancellation")
.field("cancelled", &self.is_cancelled())
.finish()
}
}
impl McpOperationCancellation {
pub(crate) fn new() -> Self {
let (sender, receiver) = oneshot::channel();
Self {
state: Arc::new(CancellationState {
cancelled: AtomicBool::new(false),
sender: Mutex::new(Some(sender)),
signal: receiver.map(|_| ()).boxed().shared(),
}),
}
}
pub(crate) fn cancel(&self) {
self.state.cancelled.store(true, Ordering::Release);
drop(
self.state
.sender
.lock()
.expect("MCP cancellation poisoned")
.take(),
);
}
pub async fn cancelled(&self) {
self.state.signal.clone().await;
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.state.cancelled.load(Ordering::Acquire)
}
}
#[derive(Clone)]
pub struct McpRequestContext<Counterpart: Role> {
server_id: McpServerAcpId,
request_id: McpRequestId,
connection: McpConnectionTo<Counterpart>,
metadata: Map<String, Value>,
cancellation: RequestCancellation,
operation_cancellation: McpOperationCancellation,
notify: Arc<Notify>,
}
impl<Counterpart: Role> std::fmt::Debug for McpRequestContext<Counterpart> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpRequestContext")
.field("server_id", &self.server_id)
.field("request_id", &self.request_id)
.field("metadata", &self.metadata)
.field("operation_cancellation", &self.operation_cancellation)
.finish_non_exhaustive()
}
}
impl<Counterpart: Role> McpRequestContext<Counterpart> {
pub(crate) fn new(
server_id: McpServerAcpId,
request_id: McpRequestId,
connection: McpConnectionTo<Counterpart>,
metadata: Map<String, Value>,
cancellation: RequestCancellation,
operation_cancellation: McpOperationCancellation,
notify: Arc<Notify>,
) -> Self {
Self {
server_id,
request_id,
connection,
metadata,
cancellation,
operation_cancellation,
notify,
}
}
pub fn server_id(&self) -> &McpServerAcpId {
&self.server_id
}
pub fn request_id(&self) -> &McpRequestId {
&self.request_id
}
pub fn connection(&self) -> &McpConnectionTo<Counterpart> {
&self.connection
}
pub fn metadata(&self) -> &Map<String, Value> {
&self.metadata
}
pub fn cancellation(&self) -> &RequestCancellation {
&self.cancellation
}
pub fn operation_cancellation(&self) -> &McpOperationCancellation {
&self.operation_cancellation
}
pub async fn send_notification(
&self,
method: impl Into<String>,
params: Option<Map<String, Value>>,
) -> Result<(), Error> {
if self.cancellation.is_cancelled() || self.operation_cancellation.is_cancelled() {
return Err(Error::request_cancelled());
}
(self.notify)(method.into(), params).await
}
}
pub trait McpService<Counterpart: Role>: Send + Sync + 'static {
fn execute(
&self,
request: McpRequest,
context: McpRequestContext<Counterpart>,
) -> BoxFuture<'static, Result<McpOutcome, Error>>;
}