use serde::{Deserialize, Serialize};
use serde_json::Value;
use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ErrorCode {
#[serde(rename = "-32700")]
ParseError = -32700,
#[serde(rename = "-32600")]
InvalidRequest = -32600,
#[serde(rename = "-32601")]
MethodNotFound = -32601,
#[serde(rename = "-32602")]
InvalidParams = -32602,
#[serde(rename = "-32603")]
InternalError = -32603,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Error)]
#[error("{message} (code: {code:?})")]
pub struct McpError {
pub code: ErrorCode,
pub message: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub data: Option<Value>,
}
impl McpError {
pub fn new(code: ErrorCode, message: impl Into<String>) -> Self {
Self {
code,
message: message.into(),
data: None,
}
}
pub fn with_data(code: ErrorCode, message: impl Into<String>, data: Value) -> Self {
Self {
code,
message: message.into(),
data: Some(data),
}
}
pub fn parse_error(message: impl Into<String>) -> Self {
Self::new(ErrorCode::ParseError, message)
}
pub fn invalid_request(message: impl Into<String>) -> Self {
Self::new(ErrorCode::InvalidRequest, message)
}
pub fn method_not_found(method: impl Into<String>) -> Self {
Self::new(
ErrorCode::MethodNotFound,
format!("Method not found: {}", method.into()),
)
}
pub fn invalid_params(message: impl Into<String>) -> Self {
Self::new(ErrorCode::InvalidParams, message)
}
pub fn internal_error(message: impl Into<String>) -> Self {
Self::new(ErrorCode::InternalError, message)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_error_serialization() {
let error = McpError::invalid_params("Missing required field");
let json = serde_json::to_string(&error).unwrap();
let deserialized: McpError = serde_json::from_str(&json).unwrap();
assert_eq!(error, deserialized);
}
#[test]
fn test_error_with_data() {
let error = McpError::with_data(
ErrorCode::InvalidParams,
"Validation failed",
json!({"field": "name", "issue": "required"}),
);
let json = serde_json::to_string(&error).unwrap();
let deserialized: McpError = serde_json::from_str(&json).unwrap();
assert_eq!(error, deserialized);
}
}