use super::binding::{CapabilityToken, McpInstallError, McpRouteTable, PendingMcpBinding};
use super::handler::TransactionMcpHandler;
use crate::transaction::dispatcher::TransactionToolDispatcher;
use crate::transaction::resolved_tools::ResolvedToolSet;
use axum::body::Body;
use axum::extract::{Path, Request, State};
use axum::http::{Response, StatusCode};
use axum::routing::any;
use axum::Router;
use monoloop_contracts::{ExchangeId, TransactionId};
use rmcp::transport::streamable_http_server::{
session::local::LocalSessionManager, StreamableHttpServerConfig, StreamableHttpService,
};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct McpGatewayHandle {
routes: Arc<McpRouteTable>,
base_url: String,
local_addr: SocketAddr,
}
impl McpGatewayHandle {
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn base_url(&self) -> &str {
&self.base_url
}
pub fn routes(&self) -> &Arc<McpRouteTable> {
&self.routes
}
pub fn install_pending(
&self,
transaction_id: TransactionId,
tools: ResolvedToolSet,
dispatcher: Arc<TransactionToolDispatcher>,
exchange_id: ExchangeId,
) -> Result<PendingMcpBinding, McpInstallError> {
self.routes.install_pending(
transaction_id,
tools,
dispatcher,
exchange_id,
&self.base_url,
)
}
pub fn activate(&self, token: &CapabilityToken) -> Result<(), McpInstallError> {
self.routes.activate(token)
}
pub fn revoke(&self, token: &CapabilityToken) -> bool {
let removed = self.routes.revoke(token);
if removed {
drop_capability_service(&token.to_hex());
}
removed
}
}
pub struct McpGateway {
handle: McpGatewayHandle,
cancel: CancellationToken,
join: JoinHandle<()>,
}
impl McpGateway {
pub async fn bind_loopback(max_routes: usize) -> Result<Self, McpInstallError> {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.map_err(|_| McpInstallError::InvalidDescriptor)?;
let local_addr = listener
.local_addr()
.map_err(|_| McpInstallError::InvalidDescriptor)?;
if !local_addr.ip().is_loopback() {
return Err(McpInstallError::InvalidDescriptor);
}
let routes = McpRouteTable::new(max_routes);
let base_url = format!("http://{}", local_addr);
let cancel = CancellationToken::new();
let cancel_serve = cancel.clone();
let routes_state = Arc::clone(&routes);
let app = Router::new()
.route("/mcp/{token}", any(mcp_dispatch))
.route("/mcp/{token}/{*rest}", any(mcp_dispatch_rest))
.with_state(routes_state);
let join = tokio::spawn(async move {
let _ = axum::serve(listener, app)
.with_graceful_shutdown(async move {
cancel_serve.cancelled().await;
})
.await;
});
Ok(Self {
handle: McpGatewayHandle {
routes,
base_url,
local_addr,
},
cancel,
join,
})
}
pub fn handle(&self) -> McpGatewayHandle {
self.handle.clone()
}
pub fn local_addr(&self) -> SocketAddr {
self.handle.local_addr()
}
pub fn base_url(&self) -> &str {
self.handle.base_url()
}
pub fn routes(&self) -> &Arc<McpRouteTable> {
self.handle.routes()
}
pub fn install_pending(
&self,
transaction_id: TransactionId,
tools: ResolvedToolSet,
dispatcher: Arc<TransactionToolDispatcher>,
exchange_id: ExchangeId,
) -> Result<PendingMcpBinding, McpInstallError> {
self.handle
.install_pending(transaction_id, tools, dispatcher, exchange_id)
}
pub fn activate(&self, token: &CapabilityToken) -> Result<(), McpInstallError> {
self.handle.activate(token)
}
pub fn revoke(&self, token: &CapabilityToken) -> bool {
self.handle.revoke(token)
}
pub async fn shutdown(self) {
let tokens = self.handle.routes.revoke_all();
for hex in tokens {
drop_capability_service(&hex);
}
self.cancel.cancel();
let _ = self.join.await;
}
}
async fn mcp_dispatch(
State(routes): State<Arc<McpRouteTable>>,
Path(token): Path<String>,
req: Request,
) -> Response<Body> {
forward_mcp(routes, &token, req).await
}
async fn mcp_dispatch_rest(
State(routes): State<Arc<McpRouteTable>>,
Path((token, _rest)): Path<(String, String)>,
req: Request,
) -> Response<Body> {
forward_mcp(routes, &token, req).await
}
struct CapabilityHttpService {
service: StreamableHttpService<TransactionMcpHandler, LocalSessionManager>,
cancel: CancellationToken,
permits: Arc<tokio::sync::Semaphore>,
}
static CAPABILITY_SERVICES: std::sync::OnceLock<
std::sync::Mutex<std::collections::HashMap<String, Arc<CapabilityHttpService>>>,
> = std::sync::OnceLock::new();
static GLOBAL_MCP_PERMITS: std::sync::OnceLock<Arc<tokio::sync::Semaphore>> =
std::sync::OnceLock::new();
const MAX_GLOBAL_MCP_REQUESTS: usize = 64;
const MAX_PER_CAPABILITY_MCP_REQUESTS: usize = 8;
const MCP_REQUEST_DURATION: std::time::Duration = std::time::Duration::from_secs(30);
fn capability_services(
) -> &'static std::sync::Mutex<std::collections::HashMap<String, Arc<CapabilityHttpService>>> {
CAPABILITY_SERVICES.get_or_init(|| std::sync::Mutex::new(std::collections::HashMap::new()))
}
fn global_mcp_permits() -> Arc<tokio::sync::Semaphore> {
GLOBAL_MCP_PERMITS
.get_or_init(|| Arc::new(tokio::sync::Semaphore::new(MAX_GLOBAL_MCP_REQUESTS)))
.clone()
}
fn drop_capability_service(token_hex: &str) {
let key = CapabilityToken::from_hex(token_hex)
.map(|t| t.to_hex())
.unwrap_or_else(|| token_hex.to_ascii_lowercase());
if let Ok(mut map) = capability_services().lock() {
if let Some(svc) = map.remove(&key) {
svc.cancel.cancel();
}
}
}
async fn forward_mcp(routes: Arc<McpRouteTable>, token_hex: &str, req: Request) -> Response<Body> {
let Some(canonical) = CapabilityToken::from_hex(token_hex).map(|t| t.to_hex()) else {
return Response::builder()
.status(StatusCode::NOT_FOUND)
.body(Body::from("unknown capability"))
.unwrap_or_else(|_| Response::new(Body::empty()));
};
let Some(binding) = routes.get_by_hex(&canonical) else {
drop_capability_service(&canonical);
return Response::builder()
.status(StatusCode::NOT_FOUND)
.body(Body::from("unknown capability"))
.unwrap_or_else(|_| Response::new(Body::empty()));
};
let (parts, body) = req.into_parts();
let collected = match axum::body::to_bytes(body, 1024 * 1024).await {
Ok(b) => b,
Err(_) => {
return Response::builder()
.status(StatusCode::PAYLOAD_TOO_LARGE)
.body(Body::from("request body exceeds bound"))
.unwrap_or_else(|_| Response::new(Body::empty()));
}
};
let req = Request::from_parts(parts, Body::from(collected));
let service = {
let mut map = capability_services()
.lock()
.unwrap_or_else(|e| e.into_inner());
map.entry(canonical.clone())
.or_insert_with(|| {
let handler = binding.handler.clone();
let cancel = CancellationToken::new();
let mut config = StreamableHttpServerConfig::default();
config.cancellation_token = cancel.clone();
config.sse_keep_alive = None;
config.sse_retry = None;
config.json_response = true;
Arc::new(CapabilityHttpService {
service: StreamableHttpService::new(
move || Ok(handler.clone()),
Arc::new(LocalSessionManager::default()),
config,
),
cancel,
permits: Arc::new(tokio::sync::Semaphore::new(MAX_PER_CAPABILITY_MCP_REQUESTS)),
})
})
.clone()
};
let Ok(_global) = global_mcp_permits().try_acquire_owned() else {
return Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.body(Body::from("mcp global concurrency exceeded"))
.unwrap_or_else(|_| Response::new(Body::empty()));
};
let Ok(_local) = service.permits.clone().try_acquire_owned() else {
return Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.body(Body::from("mcp capability concurrency exceeded"))
.unwrap_or_else(|_| Response::new(Body::empty()));
};
let req = rewrite_path(req, &canonical);
match tokio::time::timeout(MCP_REQUEST_DURATION, service.service.handle(req)).await {
Ok(response) => response.map(Body::new),
Err(_) => Response::builder()
.status(StatusCode::GATEWAY_TIMEOUT)
.body(Body::from("mcp request deadline exceeded"))
.unwrap_or_else(|_| Response::new(Body::empty())),
}
}
fn rewrite_path(req: Request, token_hex: &str) -> Request {
let (mut parts, body) = req.into_parts();
let path = parts.uri.path().to_string();
let query = parts.uri.query().map(|q| q.to_string());
let prefix = format!("/mcp/{token_hex}");
let new_path = if let Some(rest) = path.strip_prefix(&prefix) {
if rest.is_empty() {
"/".to_string()
} else {
rest.to_string()
}
} else {
path
};
let pq = match query {
Some(q) => format!("{new_path}?{q}"),
None => new_path,
};
if let Ok(uri) = pq.parse::<axum::http::Uri>() {
parts.uri = uri;
}
Request::from_parts(parts, body)
}