1use thiserror::Error;
7
8#[derive(Error, Debug, Clone)]
10pub enum McpError {
11 #[error("Transport error: {0}")]
13 Transport(String),
14
15 #[error("Protocol error: {0}")]
17 Protocol(String),
18
19 #[error("Serialization error: {0}")]
21 Serialization(String),
22
23 #[error("Invalid URI: {0}")]
25 InvalidUri(String),
26
27 #[error("Tool not found: {0}")]
29 ToolNotFound(String),
30
31 #[error("Resource not found: {0}")]
33 ResourceNotFound(String),
34
35 #[error("Prompt not found: {0}")]
37 PromptNotFound(String),
38
39 #[error("Method not found: {0}")]
41 MethodNotFound(String),
42
43 #[error("Unsupported MCP protocol version {requested}; supported: {supported:?}")]
45 UnsupportedProtocolVersion {
46 requested: String,
47 supported: Vec<String>,
48 },
49
50 #[error("MCP HTTP header mismatch: {0}")]
52 HeaderMismatch(String),
53
54 #[error("Missing required client capability: {0}")]
56 MissingRequiredClientCapability(serde_json::Value),
57
58 #[error("Invalid parameters: {0}")]
60 InvalidParams(String),
61
62 #[error("Connection error: {0}")]
64 Connection(String),
65
66 #[error("Authentication error: {0}")]
68 Authentication(String),
69
70 #[error("Authorization error: {0}")]
72 Auth(String),
73
74 #[error("Forbidden: {0}")]
76 Forbidden(String),
77
78 #[error("Rate limit exceeded; retry after {retry_after_ms}ms")]
80 RateLimited { retry_after_ms: u64 },
81
82 #[error("Validation error: {0}")]
84 Validation(String),
85
86 #[error("I/O error: {0}")]
88 Io(String),
89
90 #[error("URL error: {0}")]
92 Url(String),
93
94 #[cfg(feature = "http")]
96 #[error("HTTP error: {0}")]
97 Http(String),
98
99 #[cfg(feature = "websocket")]
101 #[error("WebSocket error: {0}")]
102 WebSocket(String),
103
104 #[error("Schema validation error: {0}")]
106 SchemaValidation(String),
107 #[error("Timeout error: {0}")]
109 Timeout(String),
110
111 #[error("Operation cancelled: {0}")]
113 Cancelled(String),
114
115 #[error("Internal error: {0}")]
117 Internal(String),
118}
119
120impl 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
139pub type McpResult<T> = Result<T, McpError>;
141
142impl McpError {
143 pub fn transport<S: Into<String>>(message: S) -> Self {
145 Self::Transport(message.into())
146 }
147
148 pub fn protocol<S: Into<String>>(message: S) -> Self {
150 Self::Protocol(message.into())
151 }
152
153 pub fn validation<S: Into<String>>(message: S) -> Self {
155 Self::Validation(message.into())
156 }
157
158 pub fn connection<S: Into<String>>(message: S) -> Self {
160 Self::Connection(message.into())
161 }
162
163 pub fn internal<S: Into<String>>(message: S) -> Self {
165 Self::Internal(message.into())
166 }
167
168 pub fn io(err: std::io::Error) -> Self {
170 Self::Io(err.to_string())
171 }
172
173 pub fn serialization(err: serde_json::Error) -> Self {
175 Self::Serialization(err.to_string())
176 }
177
178 pub fn timeout<S: Into<String>>(message: S) -> Self {
180 Self::Timeout(message.into())
181 }
182
183 pub fn connection_error<S: Into<String>>(message: S) -> Self {
185 Self::Connection(message.into())
186 }
187
188 pub fn protocol_error<S: Into<String>>(message: S) -> Self {
190 Self::Protocol(message.into())
191 }
192
193 pub fn validation_error<S: Into<String>>(message: S) -> Self {
195 Self::Validation(message.into())
196 }
197
198 pub fn timeout_error() -> Self {
200 Self::Timeout("Operation timed out".to_string())
201 }
202
203 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 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#[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#[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}