Skip to main content

prism_mcp_rs/core/
error.rs

1//! Error types for the MCP Rust SDK
2//!
3//! This module defines all error types that can occur within the MCP SDK,
4//! providing structured error handling with detailed context.
5
6use thiserror::Error;
7
8/// The main error type for the MCP SDK
9#[derive(Error, Debug, Clone)]
10pub enum McpError {
11    /// Transport-related errors (connection, I/O, etc.)
12    #[error("Transport error: {0}")]
13    Transport(String),
14
15    /// Protocol-level errors (invalid messages, unexpected responses, etc.)
16    #[error("Protocol error: {0}")]
17    Protocol(String),
18
19    /// JSON serialization/deserialization errors
20    #[error("Serialization error: {0}")]
21    Serialization(String),
22
23    /// Invalid URI format or content
24    #[error("Invalid URI: {0}")]
25    InvalidUri(String),
26
27    /// Requested tool was not found
28    #[error("Tool not found: {0}")]
29    ToolNotFound(String),
30
31    /// Requested resource was not found
32    #[error("Resource not found: {0}")]
33    ResourceNotFound(String),
34
35    /// Requested prompt was not found
36    #[error("Prompt not found: {0}")]
37    PromptNotFound(String),
38
39    /// Method not found (JSON-RPC error)
40    #[error("Method not found: {0}")]
41    MethodNotFound(String),
42
43    /// Requested MCP protocol revision is not supported by the peer.
44    #[error("Unsupported MCP protocol version {requested}; supported: {supported:?}")]
45    UnsupportedProtocolVersion {
46        requested: String,
47        supported: Vec<String>,
48    },
49
50    /// Standard MCP HTTP headers do not match the JSON-RPC body.
51    #[error("MCP HTTP header mismatch: {0}")]
52    HeaderMismatch(String),
53
54    /// A modern request omitted a capability required by the operation.
55    #[error("Missing required client capability: {0}")]
56    MissingRequiredClientCapability(serde_json::Value),
57
58    /// Invalid parameters (JSON-RPC error)
59    #[error("Invalid parameters: {0}")]
60    InvalidParams(String),
61
62    /// Connection-related errors
63    #[error("Connection error: {0}")]
64    Connection(String),
65
66    /// Authentication/authorization errors
67    #[error("Authentication error: {0}")]
68    Authentication(String),
69
70    /// OAuth 2.1 authorization errors
71    #[error("Authorization error: {0}")]
72    Auth(String),
73
74    /// The authenticated principal is not allowed to perform the operation.
75    #[error("Forbidden: {0}")]
76    Forbidden(String),
77
78    /// The caller exceeded an enforced request rate.
79    #[error("Rate limit exceeded; retry after {retry_after_ms}ms")]
80    RateLimited { retry_after_ms: u64 },
81
82    /// Input validation errors
83    #[error("Validation error: {0}")]
84    Validation(String),
85
86    /// I/O errors from the standard library
87    #[error("I/O error: {0}")]
88    Io(String),
89
90    /// URL parsing errors
91    #[error("URL error: {0}")]
92    Url(String),
93
94    /// HTTP-related errors when using HTTP transport
95    #[cfg(feature = "http")]
96    #[error("HTTP error: {0}")]
97    Http(String),
98
99    /// WebSocket-related errors when using WebSocket transport
100    #[cfg(feature = "websocket")]
101    #[error("WebSocket error: {0}")]
102    WebSocket(String),
103
104    /// JSON Schema validation errors
105    #[error("Schema validation error: {0}")]
106    SchemaValidation(String),
107    /// Timeout errors
108    #[error("Timeout error: {0}")]
109    Timeout(String),
110
111    /// Cancellation errors
112    #[error("Operation cancelled: {0}")]
113    Cancelled(String),
114
115    /// Internal errors that shouldn't normally occur
116    #[error("Internal error: {0}")]
117    Internal(String),
118}
119
120// Manual From implementations for types that don't implement Clone
121impl From<serde_json::Error> for McpError {
122    fn from(err: serde_json::Error) -> Self {
123        McpError::Serialization(err.to_string())
124    }
125}
126
127impl From<std::io::Error> for McpError {
128    fn from(err: std::io::Error) -> Self {
129        McpError::Io(err.to_string())
130    }
131}
132
133impl From<url::ParseError> for McpError {
134    fn from(err: url::ParseError) -> Self {
135        McpError::Url(err.to_string())
136    }
137}
138
139/// Result type alias for MCP operations
140pub type McpResult<T> = Result<T, McpError>;
141
142impl McpError {
143    /// Create a new transport error
144    pub fn transport<S: Into<String>>(message: S) -> Self {
145        Self::Transport(message.into())
146    }
147
148    /// Create a new protocol error
149    pub fn protocol<S: Into<String>>(message: S) -> Self {
150        Self::Protocol(message.into())
151    }
152
153    /// Create a new validation error
154    pub fn validation<S: Into<String>>(message: S) -> Self {
155        Self::Validation(message.into())
156    }
157
158    /// Create a new connection error
159    pub fn connection<S: Into<String>>(message: S) -> Self {
160        Self::Connection(message.into())
161    }
162
163    /// Create a new internal error
164    pub fn internal<S: Into<String>>(message: S) -> Self {
165        Self::Internal(message.into())
166    }
167
168    /// Create a new IO error from std::io::Error
169    pub fn io(err: std::io::Error) -> Self {
170        Self::Io(err.to_string())
171    }
172
173    /// Create a new serialization error from serde_json::Error
174    pub fn serialization(err: serde_json::Error) -> Self {
175        Self::Serialization(err.to_string())
176    }
177
178    /// Create a new timeout error
179    pub fn timeout<S: Into<String>>(message: S) -> Self {
180        Self::Timeout(message.into())
181    }
182
183    /// Create a connection error (compatibility method)
184    pub fn connection_error<S: Into<String>>(message: S) -> Self {
185        Self::Connection(message.into())
186    }
187
188    /// Create a protocol error (compatibility method)
189    pub fn protocol_error<S: Into<String>>(message: S) -> Self {
190        Self::Protocol(message.into())
191    }
192
193    /// Create a validation error (compatibility method)
194    pub fn validation_error<S: Into<String>>(message: S) -> Self {
195        Self::Validation(message.into())
196    }
197
198    /// Create a timeout error (compatibility method)
199    pub fn timeout_error() -> Self {
200        Self::Timeout("Operation timed out".to_string())
201    }
202
203    /// Check if this error is recoverable
204    pub fn is_recoverable(&self) -> bool {
205        match self {
206            McpError::Transport(_) => false,
207            McpError::Protocol(_) => false,
208            McpError::Connection(_) => true,
209            McpError::Timeout(_) => true,
210            McpError::Validation(_) => false,
211            McpError::ToolNotFound(_) => false,
212            McpError::ResourceNotFound(_) => false,
213            McpError::PromptNotFound(_) => false,
214            McpError::MethodNotFound(_) => false,
215            McpError::UnsupportedProtocolVersion { .. }
216            | McpError::HeaderMismatch(_)
217            | McpError::MissingRequiredClientCapability(_) => false,
218            McpError::InvalidParams(_) => false,
219            McpError::Authentication(_) => false,
220            McpError::Serialization(_) => false,
221            McpError::InvalidUri(_) => false,
222            McpError::Io(_) => true,
223            McpError::Url(_) => false,
224            #[cfg(feature = "http")]
225            McpError::Http(_) => true,
226            #[cfg(feature = "websocket")]
227            McpError::WebSocket(_) => true,
228            McpError::SchemaValidation(_) => false,
229            McpError::Cancelled(_) => false,
230            McpError::Auth(_) => false,
231            McpError::Forbidden(_) => false,
232            McpError::RateLimited { .. } => true,
233            McpError::Internal(_) => false,
234        }
235    }
236
237    /// Get the error category for logging/metrics
238    pub fn category(&self) -> &'static str {
239        match self {
240            McpError::Transport(_) => "transport",
241            McpError::Protocol(_) => "protocol",
242            McpError::Connection(_) => "connection",
243            McpError::Timeout(_) => "timeout",
244            McpError::Validation(_) => "validation",
245            McpError::ToolNotFound(_) => "not_found",
246            McpError::ResourceNotFound(_) => "not_found",
247            McpError::PromptNotFound(_) => "not_found",
248            McpError::MethodNotFound(_) => "not_found",
249            McpError::UnsupportedProtocolVersion { .. } => "protocol_version",
250            McpError::HeaderMismatch(_) => "protocol_header",
251            McpError::MissingRequiredClientCapability(_) => "capability",
252            McpError::InvalidParams(_) => "validation",
253            McpError::Authentication(_) => "auth",
254            McpError::Serialization(_) => "serialization",
255            McpError::InvalidUri(_) => "validation",
256            McpError::Io(_) => "io",
257            McpError::Url(_) => "validation",
258            #[cfg(feature = "http")]
259            McpError::Http(_) => "http",
260            #[cfg(feature = "websocket")]
261            McpError::WebSocket(_) => "websocket",
262            McpError::SchemaValidation(_) => "validation",
263            McpError::Cancelled(_) => "cancelled",
264            McpError::Auth(_) => "auth",
265            McpError::Forbidden(_) => "authorization",
266            McpError::RateLimited { .. } => "rate_limit",
267            McpError::Internal(_) => "internal",
268        }
269    }
270}
271
272// Convert common HTTP errors when the feature is enabled
273#[cfg(feature = "http")]
274impl From<reqwest::Error> for McpError {
275    fn from(err: reqwest::Error) -> Self {
276        McpError::Http(err.to_string())
277    }
278}
279
280// Convert common WebSocket errors when the feature is enabled
281#[cfg(feature = "websocket")]
282impl From<tokio_tungstenite::tungstenite::Error> for McpError {
283    fn from(err: tokio_tungstenite::tungstenite::Error) -> Self {
284        McpError::WebSocket(err.to_string())
285    }
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291
292    #[test]
293    fn test_error_creation() {
294        let error = McpError::transport("Connection failed");
295        assert_eq!(error.to_string(), "Transport error: Connection failed");
296        assert_eq!(error.category(), "transport");
297        assert!(!error.is_recoverable());
298    }
299
300    #[test]
301    fn test_error_recovery() {
302        assert!(McpError::connection("timeout").is_recoverable());
303        assert!(!McpError::validation("invalid input").is_recoverable());
304        assert!(McpError::timeout("request timeout").is_recoverable());
305    }
306
307    #[test]
308    fn test_error_categories() {
309        assert_eq!(McpError::protocol("bad message").category(), "protocol");
310        assert_eq!(
311            McpError::ToolNotFound("missing".to_string()).category(),
312            "not_found"
313        );
314        assert_eq!(
315            McpError::Authentication("unauthorized".to_string()).category(),
316            "auth"
317        );
318    }
319}