use crate::Node;
use crate::network::{RaftInitError, RaftSnapshotError, RaftWriteError};
use axum::Json;
use axum::http::StatusCode;
use axum::response::IntoResponse;
use bincode::error::{DecodeError, EncodeError};
use fastwebsockets::WebSocketError;
use openraft::error::{CheckIsLeaderError, ClientWriteError, Fatal, RaftError};
use serde::{Deserialize, Serialize};
use std::borrow::Cow;
use std::convert::Infallible;
use thiserror::Error;
use tokio::task::JoinError;
use tracing::trace;
#[cfg(feature = "listen_notify")]
use crate::store::state_machine::memory::notify_handler::NotifyRequest;
#[derive(Debug, Error, Serialize, Deserialize)]
pub enum Error {
#[error("BadRequest: {0}")]
BadRequest(Cow<'static, str>),
#[error("Bincode: {0}")]
Bincode(String),
#[error("Cache: {0}")]
Cache(Cow<'static, str>),
#[error("Channel: {0}")]
Channel(String),
#[error("CheckIsLeaderError: {0}")]
CheckIsLeaderError(Box<RaftError<u64, CheckIsLeaderError<u64, Node>>>),
#[error("ClientWriteError: {0}")]
ClientWriteError(Box<RaftWriteError>),
#[error("Config: {0}")]
Config(Cow<'static, str>),
#[error("Connect: {0}")]
Connect(String),
#[error("Connect: {0}")]
ConstraintViolation(String),
#[cfg(any(feature = "dashboard", feature = "s3"))]
#[error("Cryptr: {0}")]
Cryptr(String),
#[error("Error: {0}")]
Error(Cow<'static, str>),
#[error("InitializeError: {0}")]
InitializeError(Box<RaftInitError>),
#[error("LeaderChange: {0}")]
LeaderChange(Cow<'static, str>),
#[error("QueryParams: {0}")]
QueryParams(Cow<'static, str>),
#[error("QueryReturnedNoRows: {0}")]
QueryReturnedNoRows(Cow<'static, str>),
#[error("PrepareStatement: {0}")]
PrepareStatement(Cow<'static, str>),
#[error("RaftError: {0}")]
RaftError(Box<RaftError<u64>>),
#[error("RaftErrorFatal: {0}")]
RaftErrorFatal(Box<Fatal<u64>>),
#[error("Request: {0}")]
Request(String),
#[cfg(feature = "s3")]
#[error("S3: {0}")]
S3(String),
#[error("SnapshotError: {0}")]
SnapshotError(Box<RaftSnapshotError>),
#[cfg(feature = "sqlite")]
#[error("Sqlite: {0}")]
Sqlite(Cow<'static, str>),
#[error("Timeout: {0}")]
Timeout(String),
#[error("Token: {0}")]
Token(Cow<'static, str>),
#[error("Transaction: {0}")]
Transaction(Cow<'static, str>),
#[error("Unauthorized: {0}")]
Unauthorized(Cow<'static, str>),
#[error("WAL: {0}")]
WAL(String),
#[error("WebSocket: {0}")]
WebSocket(String),
#[cfg(feature = "toml")]
#[error("Error: {0}")]
String(String),
}
impl Error {
pub fn new<E: Into<Cow<'static, str>>>(error: E) -> Self {
Self::Error(error.into())
}
pub fn is_forward_to_leader(&self) -> Option<(Option<u64>, &Option<Node>)> {
if let Self::ClientWriteError(err) = self
&& let Some(err) = err.api_error()
{
match err {
ClientWriteError::ForwardToLeader(err) => {
return Some((err.leader_id, &err.leader_node));
}
ClientWriteError::ChangeMembershipError(_) => {}
}
}
if let Self::CheckIsLeaderError(err) = self
&& let Some(err) = err.forward_to_leader()
{
return Some((err.leader_id, &err.leader_node));
}
None
}
}
impl IntoResponse for Error {
fn into_response(self) -> axum::response::Response {
let status = match &self {
Error::BadRequest(_) => StatusCode::BAD_REQUEST,
Error::Bincode(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::Cache(_) => StatusCode::BAD_REQUEST,
Error::Channel(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::CheckIsLeaderError(_) => StatusCode::CONFLICT,
Error::ConstraintViolation(_) => StatusCode::BAD_REQUEST,
#[cfg(any(feature = "dashboard", feature = "s3"))]
Error::Cryptr(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::LeaderChange(_) => StatusCode::CONFLICT,
Error::QueryParams(_) => StatusCode::BAD_REQUEST,
Error::QueryReturnedNoRows(_) => StatusCode::NOT_FOUND,
Error::PrepareStatement(_) => StatusCode::BAD_REQUEST,
Error::ClientWriteError(_) => {
if self.is_forward_to_leader().is_some() {
StatusCode::PERMANENT_REDIRECT
} else {
StatusCode::INTERNAL_SERVER_ERROR
}
}
Error::Config(_) => StatusCode::BAD_REQUEST,
Error::Connect(_) => StatusCode::SERVICE_UNAVAILABLE,
Error::Error(_) => StatusCode::BAD_REQUEST,
Error::InitializeError(_) => StatusCode::BAD_REQUEST,
Error::RaftError(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::RaftErrorFatal(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::Request(_) => StatusCode::BAD_REQUEST,
#[cfg(feature = "s3")]
Error::S3(_) => StatusCode::BAD_REQUEST,
Error::SnapshotError(_) => StatusCode::INTERNAL_SERVER_ERROR,
#[cfg(feature = "sqlite")]
Error::Sqlite(_) => StatusCode::BAD_REQUEST,
Error::Timeout(_) => StatusCode::REQUEST_TIMEOUT,
Error::Token(_) => StatusCode::UNAUTHORIZED,
Error::Transaction(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::Unauthorized(_) => StatusCode::UNAUTHORIZED,
Error::WAL(_) => StatusCode::INTERNAL_SERVER_ERROR,
Error::WebSocket(_) => StatusCode::BAD_REQUEST,
#[cfg(feature = "toml")]
Error::String(_) => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, Json(self)).into_response()
}
}
impl From<std::fmt::Error> for Error {
fn from(value: std::fmt::Error) -> Self {
Self::Error(value.to_string().into())
}
}
impl From<std::io::Error> for Error {
fn from(value: std::io::Error) -> Self {
Self::Error(value.to_string().into())
}
}
impl From<Box<EncodeError>> for Error {
fn from(value: Box<EncodeError>) -> Self {
trace!("bincode::EncodeError: {value}");
Self::Bincode(value.to_string())
}
}
impl From<Box<DecodeError>> for Error {
fn from(value: Box<DecodeError>) -> Self {
trace!("bincode::DecodeError: {value}");
Self::Bincode(value.to_string())
}
}
impl From<reqwest::Error> for Error {
fn from(value: reqwest::Error) -> Self {
trace!("reqwest::Error: {value}");
if value.is_connect() {
Self::Connect(value.to_string())
} else if value.is_timeout() {
Self::Timeout(value.to_string())
} else {
Self::Request(value.to_string())
}
}
}
impl From<RaftWriteError> for Error {
fn from(value: RaftWriteError) -> Self {
trace!("ClientWriteError: {value}");
Self::ClientWriteError(Box::new(value))
}
}
impl From<RaftInitError> for Error {
fn from(value: RaftInitError) -> Self {
trace!("InitializeError: {value}");
Self::InitializeError(Box::new(value))
}
}
impl From<RaftSnapshotError> for Error {
fn from(value: RaftSnapshotError) -> Self {
trace!("SnapshotError: {value}");
Self::SnapshotError(Box::new(value))
}
}
impl From<RaftError<u64>> for Error {
fn from(value: RaftError<u64>) -> Self {
trace!("RaftError: {value}");
Self::RaftError(Box::new(value))
}
}
impl From<RaftError<u64, CheckIsLeaderError<u64, Node>>> for Error {
fn from(value: RaftError<u64, CheckIsLeaderError<u64, Node>>) -> Self {
trace!("CheckIsLeaderError: {value}");
Self::CheckIsLeaderError(Box::new(value))
}
}
impl From<Fatal<u64>> for Error {
fn from(value: Fatal<u64>) -> Self {
trace!("RaftErrorFatal: {value}");
Self::RaftErrorFatal(Box::new(value))
}
}
#[cfg(feature = "sqlite")]
impl From<rusqlite::Error> for Error {
fn from(value: rusqlite::Error) -> Self {
trace!("rusqlite::Error: {value}");
match value {
rusqlite::Error::QueryReturnedNoRows => {
Self::QueryReturnedNoRows("no rows returned".into())
}
rusqlite::Error::SqliteFailure(err, ext) => match err.code {
rusqlite::ErrorCode::ConstraintViolation => {
Self::ConstraintViolation(format!("{err} {ext:?}"))
}
_ => Self::Sqlite(format!("{err} {ext:?}").into()),
},
v => Self::Sqlite(v.to_string().into()),
}
}
}
#[cfg(feature = "sqlite")]
impl From<deadpool::unmanaged::PoolError> for Error {
fn from(value: deadpool::unmanaged::PoolError) -> Self {
trace!("Sqlite: {value}");
Self::Sqlite(value.to_string().into())
}
}
impl From<serde_json::Error> for Error {
fn from(value: serde_json::Error) -> Self {
trace!("BadRequest: {value}");
Self::BadRequest(value.to_string().into())
}
}
impl From<DecodeError> for Error {
fn from(value: DecodeError) -> Self {
trace!("DecodeError: {value}");
Self::Bincode(value.to_string())
}
}
impl From<EncodeError> for Error {
fn from(value: EncodeError) -> Self {
trace!("EncodeError: {value}");
Self::Bincode(value.to_string())
}
}
impl From<fastwebsockets::WebSocketError> for Error {
fn from(value: WebSocketError) -> Self {
trace!("WebSocket: {value}");
Self::WebSocket(value.to_string())
}
}
impl From<JoinError> for Error {
fn from(value: JoinError) -> Self {
trace!("JoinError: {value}");
Self::Error(value.to_string().into())
}
}
impl From<flume::RecvError> for Error {
fn from(value: flume::RecvError) -> Self {
trace!("flume::RecvError: {value}");
Self::Channel(value.to_string())
}
}
#[cfg(feature = "listen_notify")]
impl From<flume::SendError<NotifyRequest>> for Error {
fn from(value: flume::SendError<NotifyRequest>) -> Self {
trace!("flume::SendError<NotifyRequest>: {value}");
Self::Channel(value.to_string())
}
}
#[cfg(any(feature = "backup", feature = "s3"))]
impl From<cryptr::stream::s3::S3Error> for Error {
fn from(value: cryptr::stream::s3::S3Error) -> Self {
trace!("cryptr::stream::s3::S3Error: {value}");
Self::S3(value.to_string())
}
}
#[cfg(any(feature = "dashboard", feature = "s3"))]
impl From<cryptr::CryptrError> for Error {
fn from(value: cryptr::CryptrError) -> Self {
trace!("cryptr::CryptrError: {value}");
Self::Cryptr(value.to_string())
}
}
#[cfg(feature = "dashboard")]
impl From<argon2::password_hash::Error> for Error {
fn from(value: argon2::password_hash::Error) -> Self {
trace!("argon2::password_hash::Error: {value}");
Self::Unauthorized("invalid credentials".into())
}
}
impl From<hiqlite_wal::error::Error> for Error {
fn from(value: hiqlite_wal::error::Error) -> Self {
trace!("hiqlite_wal::error::Error: {value}");
Self::WAL(value.to_string())
}
}
impl From<Infallible> for Error {
fn from(value: Infallible) -> Self {
trace!("Infallible: {value}");
Self::Error(value.to_string().into())
}
}
#[cfg(feature = "toml")]
impl From<std::string::String> for Error {
fn from(err: std::string::String) -> Self {
Error::String(err)
}
}