use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use axum::Router;
use axum::extract::{Request, State};
use axum::http::{HeaderValue, Method, StatusCode, header};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use rmcp::transport::streamable_http_server::session::SessionManager;
use rmcp::transport::streamable_http_server::session::local::LocalSessionManager;
use rmcp::transport::streamable_http_server::{StreamableHttpServerConfig, StreamableHttpService};
use super::serve::{McpServer, ServeError};
use crate::api::Authenticator;
use crate::api::rebinding::{Rebinding, refuse_rebinding};
pub const MCP_PATH: &str = "/mcp";
const METADATA_PATH: &str = "/.well-known/oauth-protected-resource";
const UNAUTHENTICATED: &str = "this request was not authenticated";
const SESSION_HEADER: &str = "mcp-session-id";
pub const MAX_SESSIONS_PER_CALLER: usize = 8;
#[derive(Debug, Clone, Default)]
pub struct HttpConfig {
hosts: Vec<String>,
origins: Vec<String>,
protected_resource: Option<ProtectedResource>,
}
#[derive(Debug, Clone, serde::Serialize)]
struct ProtectedResource {
resource: String,
authorization_servers: Vec<String>,
bearer_methods_supported: [&'static str; 1],
}
const LOOPBACK: [&str; 3] = ["localhost", "127.0.0.1", "::1"];
impl HttpConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn allow_host(mut self, host: impl Into<String>) -> Self {
self.hosts.push(host.into());
self
}
#[must_use]
pub fn allow_origin(mut self, origin: impl Into<String>) -> Self {
self.origins.push(origin.into());
self
}
pub fn protected_resource(
mut self,
resource: impl Into<String>,
authorization_servers: Vec<String>,
) -> Result<Self, ServeError> {
if authorization_servers.iter().all(|s| s.trim().is_empty()) {
return Err(ServeError::NoAuthorizationServer);
}
self.protected_resource = Some(ProtectedResource {
resource: resource.into(),
authorization_servers,
bearer_methods_supported: ["header"],
});
Ok(self)
}
fn hosts(&self) -> Vec<String> {
LOOPBACK
.iter()
.map(|h| (*h).to_owned())
.chain(self.hosts.iter().cloned())
.collect()
}
}
#[derive(Clone)]
pub struct McpHttp {
router: Router,
stop: Arc<dyn Fn() + Send + Sync>,
}
impl std::fmt::Debug for McpHttp {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpHttp").finish_non_exhaustive()
}
}
#[derive(Clone)]
struct Guard {
auth: Arc<dyn Authenticator>,
tenant: crate::core::TenantId,
challenge: HeaderValue,
sessions: Arc<Sessions>,
}
struct Sessions {
manager: Arc<LocalSessionManager>,
owners: Mutex<HashMap<String, String>>,
}
impl Sessions {
fn owner(&self, id: &str) -> Option<String> {
self.owners
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(id)
.cloned()
}
async fn bind(&self, id: &str, actor: &str) -> bool {
let live: std::collections::HashSet<String> = self
.manager
.sessions
.read()
.await
.keys()
.map(ToString::to_string)
.collect();
let mut owners = self
.owners
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
owners.retain(|session, _| live.contains(session));
let held = owners.values().filter(|owner| *owner == actor).count();
if held >= MAX_SESSIONS_PER_CALLER {
return false;
}
owners.insert(id.to_owned(), actor.to_owned());
true
}
fn forget(&self, id: &str) {
self.owners
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(id);
}
}
impl McpHttp {
pub fn new(
server: McpServer,
auth: Arc<dyn Authenticator>,
config: &HttpConfig,
) -> Result<Self, ServeError> {
let runtime = server.runtime();
if let Some(policy) = runtime.policy() {
let problems = super::serve::policy_problems(policy.as_ref());
if !problems.is_empty() {
return Err(ServeError::PolicyUnevaluable {
problems: problems.join("; "),
});
}
}
let tenant = runtime.tenant().clone();
let server = server.authenticated();
let transport = StreamableHttpServerConfig::default()
.with_allowed_hosts(config.hosts())
.with_allowed_origins(config.origins.clone())
.enforce_origin_validation();
let token = transport.cancellation_token.clone();
let manager = Arc::new(LocalSessionManager::default());
let service =
StreamableHttpService::new(move || Ok(server.clone()), Arc::clone(&manager), transport);
let challenge = config.protected_resource.as_ref().map_or_else(
|| HeaderValue::from_static("Bearer"),
|pr| {
HeaderValue::from_str(&format!(
"Bearer resource_metadata=\"{}\"",
metadata_url(&pr.resource)
))
.unwrap_or_else(|_| HeaderValue::from_static("Bearer"))
},
);
let rebinding = Rebinding::new(config.hosts(), config.origins.clone());
let guard = Guard {
auth,
tenant,
challenge,
sessions: Arc::new(Sessions {
manager,
owners: Mutex::new(HashMap::new()),
}),
};
let mut router = Router::new()
.route_service(MCP_PATH, service)
.layer(axum::middleware::from_fn_with_state(guard, authenticate));
if let Some(document) = config.protected_resource.clone() {
let document = Arc::new(document);
let path = format!("{METADATA_PATH}{MCP_PATH}");
let metadata = axum::routing::get(move || {
let document = Arc::clone(&document);
async move { axum::Json(document.as_ref().clone()) }
});
router = router
.route(METADATA_PATH, metadata.clone())
.route(&path, metadata);
}
let router = router.layer(axum::middleware::from_fn_with_state(
rebinding,
refuse_rebinding,
));
Ok(Self {
router,
stop: Arc::new(move || token.cancel()),
})
}
pub fn router(&self) -> Router {
self.router.clone()
}
pub fn close(&self) {
(self.stop)();
}
}
fn metadata_url(resource: &str) -> String {
match resource.parse::<axum::http::Uri>() {
Ok(uri) if uri.scheme().is_some() && uri.authority().is_some() => {
let path = uri.path().trim_end_matches('/');
format!(
"{}://{}{METADATA_PATH}{path}",
uri.scheme_str().unwrap_or("https"),
uri.authority()
.map_or("", axum::http::uri::Authority::as_str)
)
}
_ => METADATA_PATH.to_owned(),
}
}
async fn authenticate(State(guard): State<Guard>, mut request: Request, next: Next) -> Response {
let caller = match guard.auth.authenticate(request.headers()).await {
Ok(caller) if caller.tenant == guard.tenant => Some(caller),
Ok(_) => {
tracing::debug!(target: "agentplane::mcp", "MCP caller is another tenant's");
None
}
Err(error) => {
tracing::warn!(target: "agentplane::mcp", %error, "MCP request not authenticated");
None
}
};
let Some(caller) = caller else {
let mut refused = (StatusCode::UNAUTHORIZED, UNAUTHENTICATED).into_response();
refused
.headers_mut()
.insert(header::WWW_AUTHENTICATE, guard.challenge.clone());
return refused;
};
let actor = caller.actor.clone();
let session = request
.headers()
.get(SESSION_HEADER)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
if let Some(id) = &session
&& guard.sessions.owner(id).is_some_and(|owner| owner != actor)
{
return (StatusCode::NOT_FOUND, "Not Found: Session not found").into_response();
}
let closing = request.method() == Method::DELETE;
request.extensions_mut().insert(caller);
let response = next.run(request).await;
match session {
Some(id) if closing && response.status().is_success() => guard.sessions.forget(&id),
Some(_) => {}
None => {
let created = response
.headers()
.get(SESSION_HEADER)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
if let Some(id) = created
&& !guard.sessions.bind(&id, &actor).await
{
let _ = guard.sessions.manager.close_session(&id.into()).await;
return (
StatusCode::TOO_MANY_REQUESTS,
"this caller holds as many MCP sessions as it may; close one first",
)
.into_response();
}
}
}
response
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::rebinding::{host_allowed, origin_allowed};
#[test]
fn no_metadata_document_names_no_authorization_server() {
assert!(matches!(
HttpConfig::new().protected_resource("https://plane.example/mcp", vec![]),
Err(ServeError::NoAuthorizationServer)
));
assert!(matches!(
HttpConfig::new().protected_resource("https://plane.example/mcp", vec![" ".into()]),
Err(ServeError::NoAuthorizationServer)
));
}
#[test]
fn the_metadata_url_inserts_the_resource_path() {
assert_eq!(
metadata_url("https://plane.example/mcp"),
"https://plane.example/.well-known/oauth-protected-resource/mcp"
);
}
#[test]
fn hosts_and_origins_compare_by_authority() {
let hosts = HttpConfig::new().allow_host("plane.example:8443").hosts();
assert!(host_allowed("127.0.0.1:8081", &hosts));
assert!(host_allowed("[::1]:8081", &hosts));
assert!(host_allowed("plane.example:8443", &hosts));
assert!(!host_allowed("plane.example:9000", &hosts));
assert!(!host_allowed("evil.example", &hosts));
let origins = ["https://app.example".to_owned()];
assert!(origin_allowed("https://app.example:443", &origins));
assert!(!origin_allowed("http://app.example", &origins));
assert!(!origin_allowed("null", &origins));
}
}