use prax_query::QueryError;
use thiserror::Error;
pub type SqlxResult<T> = Result<T, SqlxError>;
#[derive(Error, Debug)]
pub enum SqlxError {
#[error("Database error: {0}")]
Sqlx(#[from] sqlx::Error),
#[error("Configuration error: {0}")]
Config(String),
#[error("Connection error: {0}")]
Connection(String),
#[error("Query error: {0}")]
Query(String),
#[error("Deserialization error: {0}")]
Deserialization(String),
#[error("Type conversion error: {0}")]
TypeConversion(String),
#[error("Pool error: {0}")]
Pool(String),
#[error("Operation timed out after {0}ms")]
Timeout(u64),
#[error("Migration error: {0}")]
Migration(String),
#[error("Internal error: {0}")]
Internal(String),
}
impl From<SqlxError> for QueryError {
fn from(err: SqlxError) -> Self {
match err {
SqlxError::Sqlx(e) => match &e {
sqlx::Error::Io(_) | sqlx::Error::PoolTimedOut | sqlx::Error::PoolClosed => {
QueryError::connection(e.to_string())
}
sqlx::Error::RowNotFound => QueryError::not_found("row"),
sqlx::Error::Database(db_err) => match db_err.code().as_deref() {
Some("23505") | Some("23503") => {
QueryError::constraint_violation("", e.to_string())
}
Some("23502") => QueryError::invalid_input("", e.to_string()),
_ => QueryError::database(e.to_string()),
},
_ => QueryError::database(e.to_string()),
},
SqlxError::Config(msg) => QueryError::connection(msg),
SqlxError::Connection(msg) => QueryError::connection(msg),
SqlxError::Query(msg) => QueryError::database(msg),
SqlxError::Deserialization(msg) => QueryError::serialization(msg),
SqlxError::TypeConversion(msg) => QueryError::serialization(msg),
SqlxError::Pool(msg) => QueryError::connection(msg),
SqlxError::Timeout(ms) => QueryError::timeout(ms),
SqlxError::Migration(msg) => QueryError::database(msg),
SqlxError::Internal(msg) => QueryError::internal(msg),
}
}
}
impl SqlxError {
pub fn config(msg: impl Into<String>) -> Self {
Self::Config(msg.into())
}
pub fn connection(msg: impl Into<String>) -> Self {
Self::Connection(msg.into())
}
pub fn query(msg: impl Into<String>) -> Self {
Self::Query(msg.into())
}
pub fn pool(msg: impl Into<String>) -> Self {
Self::Pool(msg.into())
}
pub fn timeout(ms: u64) -> Self {
Self::Timeout(ms)
}
pub fn type_conversion(msg: impl Into<String>) -> Self {
Self::TypeConversion(msg.into())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_error_creation() {
let err = SqlxError::config("test config error");
assert!(matches!(err, SqlxError::Config(_)));
let err = SqlxError::connection("test connection error");
assert!(matches!(err, SqlxError::Connection(_)));
let err = SqlxError::timeout(5000);
assert!(matches!(err, SqlxError::Timeout(5000)));
}
#[test]
fn test_error_to_query_error() {
let err = SqlxError::connection("connection failed");
let query_err: QueryError = err.into();
assert!(query_err.to_string().contains("connection"));
let err = SqlxError::timeout(1000);
let query_err: QueryError = err.into();
assert!(
query_err.to_string().contains("timeout") || query_err.to_string().contains("1000")
);
}
#[derive(Debug)]
struct MockDbError {
message: String,
code: Option<String>,
}
impl std::fmt::Display for MockDbError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.message)
}
}
impl std::error::Error for MockDbError {}
impl sqlx::error::DatabaseError for MockDbError {
fn message(&self) -> &str {
&self.message
}
fn code(&self) -> Option<std::borrow::Cow<'_, str>> {
self.code.as_deref().map(std::borrow::Cow::Borrowed)
}
fn as_error(&self) -> &(dyn std::error::Error + Send + Sync + 'static) {
self
}
fn as_error_mut(&mut self) -> &mut (dyn std::error::Error + Send + Sync + 'static) {
self
}
fn into_error(self: Box<Self>) -> Box<dyn std::error::Error + Send + Sync + 'static> {
self
}
fn kind(&self) -> sqlx::error::ErrorKind {
sqlx::error::ErrorKind::Other
}
}
#[test]
fn test_sqlx_database_error_code_mapping() {
let db_err = MockDbError {
message: "db says no".to_string(),
code: Some("23505".to_string()),
};
let err = SqlxError::from(sqlx::Error::Database(Box::new(db_err)));
let query_err: QueryError = err.into();
assert_eq!(query_err.code, prax_query::ErrorCode::UniqueConstraint);
let db_err = MockDbError {
message: "db says no".to_string(),
code: Some("23502".to_string()),
};
let err = SqlxError::from(sqlx::Error::Database(Box::new(db_err)));
let query_err: QueryError = err.into();
assert_eq!(query_err.code, prax_query::ErrorCode::InvalidParameter);
let db_err = MockDbError {
message: "deadlock detected".to_string(),
code: Some("40P01".to_string()),
};
let err = SqlxError::from(sqlx::Error::Database(Box::new(db_err)));
let query_err: QueryError = err.into();
assert_eq!(query_err.code, prax_query::ErrorCode::DatabaseError);
}
#[test]
fn test_sqlx_io_error_maps_to_connection() {
let io_err = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused");
let err = SqlxError::from(sqlx::Error::Io(io_err));
let query_err: QueryError = err.into();
assert_eq!(query_err.code, prax_query::ErrorCode::ConnectionFailed);
}
}