use std::sync::Arc;
use axum::{
Router,
body::{Body, Bytes},
http::{HeaderMap, Method, StatusCode, header::HeaderName},
response::{IntoResponse, Response},
routing::any,
};
use crate::ServerState;
use crate::api::http::auth::HttpCaller;
use super::runtime::McpRuntime;
pub(crate) const MCP_PATH: &str = "/mcp";
pub(crate) fn mcp_router(runtime: Arc<McpRuntime>) -> Router<ServerState> {
Router::new()
.route(MCP_PATH, any(handle))
.layer(axum::Extension(runtime))
}
pub(crate) fn mcp_disabled_router() -> Router<ServerState> {
Router::new().route(MCP_PATH, any(|| async { StatusCode::NOT_FOUND }))
}
async fn handle(
axum::Extension(runtime): axum::Extension<Arc<McpRuntime>>,
HttpCaller(caller): HttpCaller,
method: Method,
headers: HeaderMap,
body: Bytes,
) -> Response {
let request = aion_mcp::HttpRequest::new(
method.as_str().to_ascii_uppercase(),
aion_mcp::protocol::headers::HeaderView::new(header_pairs(&headers)),
body.to_vec(),
);
let response = runtime.server().handle_http(&request, &caller).await;
into_axum(response)
}
fn header_pairs(headers: &HeaderMap) -> Vec<(String, String)> {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_owned(), value.to_owned()))
})
.collect()
}
fn into_axum(response: aion_mcp::HttpResponse) -> Response {
let status = StatusCode::from_u16(response.status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
let mut builder = Response::builder().status(status);
for (name, value) in response.headers {
match HeaderName::from_bytes(name.as_bytes()) {
Ok(name) => builder = builder.header(name, value),
Err(error) => {
tracing::error!(
target: "aion_server::mcp",
header = %name,
"an MCP response header name was not valid and was dropped: {error}"
);
}
}
}
builder
.body(Body::from(response.body))
.unwrap_or_else(|error| {
tracing::error!(
target: "aion_server::mcp",
"an MCP response could not be assembled: {error}"
);
StatusCode::INTERNAL_SERVER_ERROR.into_response()
})
}
#[cfg(test)]
mod tests {
use axum::http::{HeaderMap, HeaderValue, header::HeaderName};
use super::header_pairs;
#[test]
fn a_non_utf8_header_value_is_dropped_not_lossily_decoded()
-> Result<(), Box<dyn std::error::Error>> {
let mut headers = HeaderMap::new();
drop(headers.insert(
HeaderName::from_static("mcp-method"),
HeaderValue::from_static("tools/list"),
));
drop(headers.insert(
HeaderName::from_static("mcp-name"),
HeaderValue::from_bytes(&[0xff, 0xfe])?,
));
let pairs = header_pairs(&headers);
assert_eq!(pairs.len(), 1);
assert_eq!(pairs[0].0, "mcp-method");
Ok(())
}
}