//! Error handling for the serving API use axum::{ Json, http::StatusCode, response::{IntoResponse, Response}, }; use serde::{Deserialize, Serialize}; use thiserror::Error; // TODO: Re-enable when rtx-inference is fixed // use rtx_inference::InferenceError; // Temporary placeholder for InferenceError until rtx-inference is fixed #[derive(Debug, Error)] pub enum InferenceError { #[error("Model not found: {0}")] ModelNotFound(String), #[error("Invalid input: {0}")] InvalidInput(String), #[error("Runtime error: {0}")] RuntimeError(String), } /// API result type alias pub type ApiResult = Result; /// API error types #[derive(Debug, Error)] pub enum ApiError { /// Inference runtime error #[error("Inference error: {0}")] Inference(#[from] InferenceError), /// Validation error #[error("Validation error: {message}")] Validation { /// Error message message: String, }, /// Not found error #[error("Not found: {resource}")] NotFound { /// Resource that was not found resource: String, }, /// Internal server error #[error("Internal server error: {message}")] Internal { /// Error message message: String, }, /// Not implemented error #[error("Not implemented: {feature}")] NotImplemented { /// Feature that is not implemented feature: String, }, /// Bad request error #[error("Bad request: {message}")] BadRequest { /// Error message message: String, }, /// Service unavailable error #[error("Service unavailable: {reason}")] ServiceUnavailable { /// Reason for unavailability reason: String, }, /// WebSocket error #[error("WebSocket error: {0}")] WebSocket(String), /// Load balancer error #[error("Load balancer error: {0}")] LoadBalancer(String), /// Serialization error #[error("Serialization error: {0}")] Serialization(#[from] serde_json::Error), } impl ApiError { /// Create a validation error pub fn validation>(message: S) -> Self { Self::Validation { message: message.into(), } } /// Create a not found error pub fn not_found>(resource: S) -> Self { Self::NotFound { resource: resource.into(), } } /// Create an internal error pub fn internal>(message: S) -> Self { Self::Internal { message: message.into(), } } /// Create a not implemented error pub fn not_implemented>(feature: S) -> Self { Self::NotImplemented { feature: feature.into(), } } /// Create a bad request error pub fn bad_request>(message: S) -> Self { Self::BadRequest { message: message.into(), } } /// Create a service unavailable error pub fn service_unavailable>(reason: S) -> Self { Self::ServiceUnavailable { reason: reason.into(), } } /// Get the HTTP status code for this error #[must_use] pub fn status_code(&self) -> StatusCode { match self { Self::Validation { .. } | Self::BadRequest { .. } => StatusCode::BAD_REQUEST, Self::NotFound { .. } => StatusCode::NOT_FOUND, Self::NotImplemented { .. } => StatusCode::NOT_IMPLEMENTED, Self::ServiceUnavailable { .. } => StatusCode::SERVICE_UNAVAILABLE, Self::Internal { .. } | Self::Inference { .. } | Self::WebSocket(_) | Self::LoadBalancer(_) | Self::Serialization(_) => StatusCode::INTERNAL_SERVER_ERROR, } } } /// Error response payload #[derive(Debug, Serialize, Deserialize)] pub struct ErrorResponse { /// Error code pub error: String, /// Human-readable message pub message: String, /// Additional details (optional) pub details: Option, /// Timestamp pub timestamp: chrono::DateTime, } impl IntoResponse for ApiError { fn into_response(self) -> Response { let status = self.status_code(); let error_response = ErrorResponse { error: match &self { Self::Validation { .. } => "validation_error", Self::NotFound { .. } => "not_found", Self::Internal { .. } => "internal_error", Self::NotImplemented { .. } => "not_implemented", Self::BadRequest { .. } => "bad_request", Self::ServiceUnavailable { .. } => "service_unavailable", Self::Inference { .. } => "inference_error", Self::WebSocket(_) => "websocket_error", Self::LoadBalancer(_) => "load_balancer_error", Self::Serialization(_) => "serialization_error", } .to_string(), message: self.to_string(), details: None, timestamp: chrono::Utc::now(), }; (status, Json(error_response)).into_response() } } #[cfg(test)] mod tests { use super::*; #[test] fn test_error_status_codes() { assert_eq!( ApiError::validation("test").status_code(), StatusCode::BAD_REQUEST ); assert_eq!( ApiError::not_found("test").status_code(), StatusCode::NOT_FOUND ); assert_eq!( ApiError::internal("test").status_code(), StatusCode::INTERNAL_SERVER_ERROR ); assert_eq!( ApiError::not_implemented("test").status_code(), StatusCode::NOT_IMPLEMENTED ); assert_eq!( ApiError::service_unavailable("test").status_code(), StatusCode::SERVICE_UNAVAILABLE ); } #[test] fn test_error_response_serialization() { let error = ApiError::validation("Invalid input"); let response = ErrorResponse { error: "validation_error".to_string(), message: error.to_string(), details: None, timestamp: chrono::Utc::now(), }; let json = serde_json::to_string(&response).unwrap(); assert!(json.contains("validation_error")); assert!(json.contains("Invalid input")); } }