use std::net::AddrParseError;
use axum::{http::StatusCode, response::IntoResponse};
use thiserror::Error;
use rust_mcp_sdk::mcp_http::McpHttpError;
use rust_mcp_sdk::auth::AuthenticationError;
pub type TransportServerResult<T> = core::result::Result<T, TransportServerError>;
#[derive(Debug, Error, Clone)]
pub enum TransportServerError {
#[error("'sessionId' query string is missing!")]
SessionIdMissing,
#[error("No session found for the given ID: {0}.")]
SessionIdInvalid(String),
#[error("Stream IO Error: {0}.")]
StreamIoError(String),
#[error("{0}")]
AddrParseError(#[from] AddrParseError),
#[error("{0}")]
HttpError(String),
#[error("Server start error: {0}")]
ServerStartError(String),
#[error("Invalid options: {0}")]
InvalidServerOptions(String),
#[error("{0}")]
SslCertError(String),
#[error("{0}")]
TransportError(String),
#[error("{0}")]
AuthenticationError(#[from] AuthenticationError),
}
impl IntoResponse for TransportServerError {
fn into_response(self) -> axum::response::Response {
let mut response = StatusCode::INTERNAL_SERVER_ERROR.into_response();
response.extensions_mut().insert(self);
response
}
}
impl From<McpHttpError> for TransportServerError {
fn from(err: McpHttpError) -> Self {
match err {
McpHttpError::StreamIoError(s) => TransportServerError::StreamIoError(s),
McpHttpError::HttpError(s) => TransportServerError::HttpError(s),
McpHttpError::TransportError(s) => TransportServerError::TransportError(s),
}
}
}
impl From<TransportServerError> for McpHttpError {
fn from(err: TransportServerError) -> Self {
match err {
TransportServerError::SessionIdMissing => {
McpHttpError::HttpError("'sessionId' query string is missing!".into())
}
TransportServerError::SessionIdInvalid(s) => {
McpHttpError::HttpError(format!("No session found for the given ID: {}.", s))
}
TransportServerError::StreamIoError(s) => McpHttpError::StreamIoError(s),
TransportServerError::HttpError(s) => McpHttpError::HttpError(s),
TransportServerError::TransportError(s) => McpHttpError::TransportError(s),
TransportServerError::AuthenticationError(e) => McpHttpError::HttpError(e.to_string()),
TransportServerError::AddrParseError(e) => McpHttpError::HttpError(e.to_string()),
TransportServerError::ServerStartError(s) => McpHttpError::HttpError(s),
TransportServerError::InvalidServerOptions(s) => McpHttpError::HttpError(s),
TransportServerError::SslCertError(s) => McpHttpError::HttpError(s),
}
}
}
impl From<TransportServerError> for rust_mcp_sdk::error::McpSdkError {
fn from(err: TransportServerError) -> Self {
rust_mcp_sdk::error::McpSdkError::Internal {
description: err.to_string(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use rust_mcp_sdk::mcp_http::{McpHttpError, McpHttpResult};
#[test]
fn mcp_to_transport_stream_io_error() {
let m = McpHttpError::StreamIoError("io".into());
let t: TransportServerError = m.into();
assert!(matches!(t, TransportServerError::StreamIoError(ref s) if s == "io"));
}
#[test]
fn mcp_to_transport_http_error() {
let m = McpHttpError::HttpError("fail".into());
let t: TransportServerError = m.into();
assert!(matches!(t, TransportServerError::HttpError(ref s) if s == "fail"));
}
#[test]
fn mcp_to_transport_transport_error() {
let m = McpHttpError::TransportError("tcp".into());
let t: TransportServerError = m.into();
assert!(matches!(t, TransportServerError::TransportError(ref s) if s == "tcp"));
}
#[test]
fn transport_to_mcp_session_id_missing() {
let t = TransportServerError::SessionIdMissing;
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(_)));
}
#[test]
fn transport_to_mcp_session_id_invalid() {
let t = TransportServerError::SessionIdInvalid("s2".into());
let m: McpHttpError = t.into();
assert!(
matches!(m, McpHttpError::HttpError(ref s) if s == "No session found for the given ID: s2.")
);
}
#[test]
fn transport_to_mcp_stream_io_error() {
let t = TransportServerError::StreamIoError("eof".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::StreamIoError(ref s) if s == "eof"));
}
#[test]
fn transport_to_mcp_http_error() {
let t = TransportServerError::HttpError("gone".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if s == "gone"));
}
#[test]
fn transport_to_mcp_transport_error() {
let t = TransportServerError::TransportError("tls".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::TransportError(ref s) if s == "tls"));
}
#[test]
fn transport_to_mcp_addr_parse_lossy() {
use std::net::AddrParseError;
let parse_err: AddrParseError = ":::".parse::<std::net::IpAddr>().unwrap_err();
let t = TransportServerError::AddrParseError(parse_err);
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if !s.is_empty()));
}
#[test]
fn transport_to_mcp_server_start_lossy() {
let t = TransportServerError::ServerStartError("port in use".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if s == "port in use"));
}
#[test]
fn transport_to_mcp_invalid_options_lossy() {
let t = TransportServerError::InvalidServerOptions("bad config".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if s == "bad config"));
}
#[test]
fn transport_to_mcp_ssl_cert_lossy() {
let t = TransportServerError::SslCertError("cert expired".into());
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if s == "cert expired"));
}
#[test]
fn transport_to_mcp_authentication_lossy() {
let auth_err = AuthenticationError::InactiveToken;
let t = TransportServerError::AuthenticationError(auth_err);
let m: McpHttpError = t.into();
assert!(matches!(m, McpHttpError::HttpError(ref s) if s.contains("Inactive")));
}
#[test]
fn round_trip_stream_io_error() {
let m = McpHttpError::StreamIoError("pipe".into());
let t: TransportServerError = m.clone().into();
let back: McpHttpError = t.into();
assert_eq!(format!("{}", m), format!("{}", back));
}
#[test]
fn round_trip_http_error() {
let m = McpHttpError::HttpError("round".into());
let t: TransportServerError = m.clone().into();
let back: McpHttpError = t.into();
assert_eq!(format!("{}", m), format!("{}", back));
}
#[test]
fn round_trip_transport_error() {
let m = McpHttpError::TransportError("round".into());
let t: TransportServerError = m.clone().into();
let back: McpHttpError = t.into();
assert_eq!(format!("{}", m), format!("{}", back));
}
#[test]
fn reverse_round_trip_session_id_missing() {
let t = TransportServerError::SessionIdMissing;
let m: McpHttpError = t.clone().into();
let back: TransportServerError = m.into();
assert_eq!(format!("{}", t), format!("{}", back));
}
#[test]
fn reverse_round_trip_session_id_invalid() {
let t = TransportServerError::SessionIdInvalid("rev".into());
let m: McpHttpError = t.clone().into();
let back: TransportServerError = m.into();
assert_eq!(format!("{}", t), format!("{}", back));
}
#[test]
fn reverse_round_trip_stream_io_error() {
let t = TransportServerError::StreamIoError("rev".into());
let m: McpHttpError = t.clone().into();
let back: TransportServerError = m.into();
assert_eq!(format!("{}", t), format!("{}", back));
}
#[test]
fn reverse_round_trip_http_error() {
let t = TransportServerError::HttpError("rev".into());
let m: McpHttpError = t.clone().into();
let back: TransportServerError = m.into();
assert_eq!(format!("{}", t), format!("{}", back));
}
#[test]
fn reverse_round_trip_transport_error() {
let t = TransportServerError::TransportError("rev".into());
let m: McpHttpError = t.clone().into();
let back: TransportServerError = m.into();
assert_eq!(format!("{}", t), format!("{}", back));
}
#[test]
fn mcp_http_result_from_transport_error() {
let r: TransportServerResult<()> = Err(TransportServerError::SessionIdInvalid("x".into()));
let m: McpHttpResult<()> = r.map_err(Into::into);
assert!(matches!(m.unwrap_err(), McpHttpError::HttpError(_)));
}
}