use crate::wire::config::ErrorConvention;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ClientError {
#[error("auth error: {message}")]
Auth {
message: String,
},
#[error("server error: {message}")]
Server {
message: String,
code: Option<String>,
},
#[error("connection error: {message}")]
Connection {
message: String,
},
#[error("timed out")]
Timeout,
#[error("frame too large: {message}")]
FrameTooLarge {
message: String,
},
#[error("decode error: {message}")]
Decode {
message: String,
},
}
impl ClientError {
pub fn from_server_message(message: impl Into<String>, convention: ErrorConvention) -> Self {
let message = message.into();
match convention {
ErrorConvention::None => Self::Server {
message,
code: None,
},
ErrorConvention::Resp3Prefixes => {
if starts_with_auth_prefix(&message) {
Self::Auth { message }
} else {
Self::Server {
message,
code: None,
}
}
}
ErrorConvention::BracketCode | ErrorConvention::Both => {
let (code, rest) = split_bracket_code(&message);
if starts_with_auth_prefix(rest) {
Self::Auth { message }
} else {
Self::Server { message, code }
}
}
}
}
}
fn starts_with_auth_prefix(message: &str) -> bool {
["NOAUTH", "WRONGPASS", "NOPERM"].iter().any(|prefix| {
message
.strip_prefix(prefix)
.is_some_and(|rest| rest.is_empty() || rest.starts_with(' '))
})
}
fn split_bracket_code(message: &str) -> (Option<String>, &str) {
if let Some(inner) = message.strip_prefix('[') {
if let Some(end) = inner.find(']') {
let code = &inner[..end];
let after = &inner[end + 1..];
if !code.is_empty() && !code.contains(char::is_whitespace) {
if let Some(rest) = after.strip_prefix(' ') {
return (Some(code.to_owned()), rest);
}
}
}
}
(None, message)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn resp3_auth_prefixes_map_to_auth_class() {
for msg in [
"NOAUTH Authentication required.",
"WRONGPASS invalid username-password pair or user is disabled.",
"NOPERM this user has no permissions",
"NOAUTH",
] {
let err = ClientError::from_server_message(msg, ErrorConvention::Resp3Prefixes);
assert_eq!(
err,
ClientError::Auth {
message: msg.to_owned()
},
"{msg} must map to the auth class (CLT-051)"
);
}
}
#[test]
fn resp3_err_prefix_is_generic_server_error_without_code() {
let err =
ClientError::from_server_message("ERR unknown command", ErrorConvention::Resp3Prefixes);
assert_eq!(
err,
ClientError::Server {
message: "ERR unknown command".to_owned(),
code: None,
}
);
}
#[test]
fn resp3_prefix_must_be_word_aligned() {
let err = ClientError::from_server_message("NOAUTHx nope", ErrorConvention::Resp3Prefixes);
assert!(matches!(err, ClientError::Server { .. }));
}
#[test]
fn bracket_code_extracts_structured_code_and_keeps_raw_message() {
let raw = "[collection_not_found] no such collection: docs";
let err = ClientError::from_server_message(raw, ErrorConvention::BracketCode);
assert_eq!(
err,
ClientError::Server {
message: raw.to_owned(),
code: Some("collection_not_found".to_owned()),
}
);
}
#[test]
fn bracket_code_still_maps_auth_prefixes_to_auth_class() {
let raw = "[unauthorized] NOAUTH token expired";
let err = ClientError::from_server_message(raw, ErrorConvention::BracketCode);
assert_eq!(
err,
ClientError::Auth {
message: raw.to_owned()
}
);
}
#[test]
fn both_convention_composes_bracket_and_prefixes() {
let err = ClientError::from_server_message(
"[wrongpass] WRONGPASS bad credentials",
ErrorConvention::Both,
);
assert!(matches!(err, ClientError::Auth { .. }));
let err = ClientError::from_server_message(
"[index_missing] ERR no such index",
ErrorConvention::Both,
);
assert_eq!(
err,
ClientError::Server {
message: "[index_missing] ERR no such index".to_owned(),
code: Some("index_missing".to_owned()),
}
);
}
#[test]
fn none_convention_never_parses() {
let err = ClientError::from_server_message("NOAUTH raw passthrough", ErrorConvention::None);
assert_eq!(
err,
ClientError::Server {
message: "NOAUTH raw passthrough".to_owned(),
code: None,
}
);
}
#[test]
fn malformed_bracket_prefixes_are_left_alone() {
for msg in ["[] empty", "[has space] x", "[nospace]tail", "[unclosed"] {
let err = ClientError::from_server_message(msg, ErrorConvention::BracketCode);
assert_eq!(
err,
ClientError::Server {
message: msg.to_owned(),
code: None,
},
"{msg} must not yield a code"
);
}
}
}