1use std::fmt;
2
3pub struct DbError(sqlx::Error);
17
18impl DbError {
19 pub fn is_unique_violation(&self) -> bool {
21 matches!(&self.0, sqlx::Error::Database(db) if db.is_unique_violation())
22 }
23
24 pub fn is_foreign_key_violation(&self) -> bool {
26 matches!(&self.0, sqlx::Error::Database(db) if db.is_foreign_key_violation())
27 }
28
29 pub fn is_row_not_found(&self) -> bool {
31 matches!(self.0, sqlx::Error::RowNotFound)
32 }
33
34 pub fn is_retryable(&self) -> bool {
38 let sqlx::Error::Database(db) = &self.0 else {
39 return false;
40 };
41 match db.code() {
42 Some(code) if code == "40001" || code == "40P01" => true,
43 Some(code) => code.parse::<i32>().is_ok_and(|c| matches!(c & 0xff, 5 | 6)),
45 None => false,
46 }
47 }
48
49 pub fn is_timeout(&self) -> bool {
51 matches!(self.0, sqlx::Error::PoolTimedOut)
52 }
53
54 pub fn sqlx(&self) -> &sqlx::Error {
57 &self.0
58 }
59}
60
61impl From<sqlx::Error> for DbError {
62 fn from(err: sqlx::Error) -> Self {
63 Self(err)
64 }
65}
66
67impl fmt::Display for DbError {
68 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
69 fmt::Display::fmt(&self.0, f)
70 }
71}
72
73impl fmt::Debug for DbError {
74 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
75 fmt::Debug::fmt(&self.0, f)
76 }
77}
78
79impl std::error::Error for DbError {
80 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
81 self.0.source()
83 }
84}
85
86#[cfg(test)]
88pub(crate) mod fake {
89 use std::borrow::Cow;
90 use std::fmt;
91
92 use sqlx::error::{DatabaseError, ErrorKind};
93
94 #[derive(Debug)]
95 struct Coded(Option<&'static str>, ErrorKind);
96
97 impl fmt::Display for Coded {
98 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
99 write!(f, "error {:?}", self.0)
100 }
101 }
102
103 impl std::error::Error for Coded {}
104
105 impl DatabaseError for Coded {
106 fn message(&self) -> &str {
107 "fake"
108 }
109 fn code(&self) -> Option<Cow<'_, str>> {
110 self.0.map(Cow::Borrowed)
111 }
112 fn as_error(&self) -> &(dyn std::error::Error + Send + Sync + 'static) {
113 self
114 }
115 fn as_error_mut(&mut self) -> &mut (dyn std::error::Error + Send + Sync + 'static) {
116 self
117 }
118 fn into_error(self: Box<Self>) -> Box<dyn std::error::Error + Send + Sync + 'static> {
119 self
120 }
121 fn kind(&self) -> ErrorKind {
122 match self.1 {
123 ErrorKind::UniqueViolation => ErrorKind::UniqueViolation,
124 _ => ErrorKind::Other,
125 }
126 }
127 }
128
129 pub(crate) fn coded(code: Option<&'static str>) -> sqlx::Error {
131 sqlx::Error::Database(Box::new(Coded(code, ErrorKind::Other)))
132 }
133
134 pub(crate) fn unique() -> sqlx::Error {
136 sqlx::Error::Database(Box::new(Coded(Some("23505"), ErrorKind::UniqueViolation)))
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::{DbError, fake};
143
144 #[test]
147 fn retryable_codes() {
148 for code in ["40001", "40P01", "5", "517", "6", "262"] {
149 assert!(
150 DbError::from(fake::coded(Some(code))).is_retryable(),
151 "{code}"
152 );
153 }
154 for code in [Some("23505"), Some("19"), Some("XX000"), None] {
155 assert!(!DbError::from(fake::coded(code)).is_retryable(), "{code:?}");
156 }
157 assert!(!DbError::from(sqlx::Error::PoolTimedOut).is_retryable());
158 assert!(DbError::from(sqlx::Error::PoolTimedOut).is_timeout());
159 assert!(!DbError::from(sqlx::Error::PoolClosed).is_unique_violation());
160 }
161
162 #[test]
165 fn unique_violations_through_every_wrapper() {
166 assert!(DbError::from(fake::unique()).is_unique_violation());
167 assert!(crate::Error::from(DbError::from(fake::unique())).is_unique_violation());
168 let raw = crate::Error::from(anyhow::Error::from(fake::unique()));
169 assert!(raw.is_unique_violation());
170 assert!(!crate::Error::from(anyhow::anyhow!("no")).is_unique_violation());
171 assert!(!crate::Error::NotFound.is_unique_violation());
172 }
173}