use std::fmt;
use std::time::Duration;
pub(crate) const ACQUIRE_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionTrouble {
Lost,
TimedOut,
Closed,
}
impl fmt::Display for ConnectionTrouble {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ConnectionTrouble::Lost => {
f.write_str("The connection to the server was lost. Reconnect to open a fresh one.")
}
ConnectionTrouble::TimedOut => write!(
f,
"No connection to the server came free within {} seconds: every connection \
is busy with a running query, or the server has stopped answering. Cancel a \
running query, or Reconnect.",
ACQUIRE_TIMEOUT.as_secs()
),
ConnectionTrouble::Closed => {
f.write_str("This connection was closed. Reconnect to open a fresh one.")
}
}
}
}
impl std::error::Error for ConnectionTrouble {}
impl ConnectionTrouble {
pub fn of(error: &anyhow::Error) -> Option<Self> {
error.chain().find_map(|cause| {
if let Some(trouble) = cause.downcast_ref::<ConnectionTrouble>() {
return Some(*trouble);
}
cause.downcast_ref::<sqlx::Error>().and_then(classify)
})
}
}
pub(crate) fn plain(error: anyhow::Error) -> anyhow::Error {
match ConnectionTrouble::of(&error) {
Some(trouble) => anyhow::Error::new(trouble),
None => error,
}
}
const MYSQL_IDLE_DISCONNECT: u16 = 4031;
fn classify(error: &sqlx::Error) -> Option<ConnectionTrouble> {
match error {
sqlx::Error::Io(_) | sqlx::Error::WorkerCrashed => Some(ConnectionTrouble::Lost),
sqlx::Error::PoolTimedOut => Some(ConnectionTrouble::TimedOut),
sqlx::Error::PoolClosed => Some(ConnectionTrouble::Closed),
sqlx::Error::Database(error) => {
let mysql = error
.try_downcast_ref::<sqlx::mysql::MySqlDatabaseError>()
.map(|error| error.number());
(mysql == Some(MYSQL_IDLE_DISCONNECT)
|| matches!(
error.code().as_deref(),
Some("57P01" | "57P02" | "57P03" | "57P04" | "57P05" | "25P03")
) && mysql.is_none())
.then_some(ConnectionTrouble::Lost)
}
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn io(kind: std::io::ErrorKind) -> anyhow::Error {
anyhow::Error::new(sqlx::Error::Io(std::io::Error::new(kind, "driver text")))
}
#[test]
fn a_dropped_socket_reads_as_a_lost_connection() {
for kind in [
std::io::ErrorKind::ConnectionReset,
std::io::ErrorKind::BrokenPipe,
std::io::ErrorKind::UnexpectedEof,
] {
let error = plain(io(kind));
assert_eq!(ConnectionTrouble::of(&error), Some(ConnectionTrouble::Lost));
let shown = format!("{error:#}");
assert!(
shown.contains("connection to the server was lost"),
"{shown}"
);
assert!(!shown.contains("driver text"), "{shown}");
}
}
#[test]
fn a_pool_timeout_names_the_wait_and_what_to_do() {
let shown = format!("{:#}", plain(sqlx::Error::PoolTimedOut.into()));
assert!(shown.contains("10 seconds"), "{shown}");
assert!(shown.contains("Reconnect"), "{shown}");
}
#[test]
fn trouble_is_found_under_added_context() {
let error = io(std::io::ErrorKind::ConnectionReset).context("statement 2 of 3");
assert_eq!(ConnectionTrouble::of(&error), Some(ConnectionTrouble::Lost));
}
#[test]
fn an_ordinary_error_is_left_alone() {
let error = plain(anyhow::anyhow!("syntax error at or near \"SELEC\""));
assert_eq!(ConnectionTrouble::of(&error), None);
assert_eq!(format!("{error:#}"), "syntax error at or near \"SELEC\"");
let error = plain(sqlx::Error::RowNotFound.into());
assert_eq!(ConnectionTrouble::of(&error), None);
}
}