pub use http::StatusCode;
use s2_api::v1 as api;
pub use s2_api::v1::error::ErrorCode;
pub use crate::session::{
append::AppendSessionError,
read::{CaughtUpError, ReadSessionError},
};
use crate::{
api::{ApiError, ServerErrorBody},
client,
types::{FencingToken, StreamPosition, ValidationError},
};
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum ClientError {
#[error("connect: {0}")]
Connect(String),
#[error("timeout")]
Timeout,
#[error("connection closed early: {0}")]
ConnectionClosedEarly(String),
#[error("request canceled: {0}")]
RequestCanceled(String),
#[error("unexpected eof: {0}")]
UnexpectedEof(String),
#[error("connection reset: {0}")]
ConnectionReset(String),
#[error("connection aborted: {0}")]
ConnectionAborted(String),
#[error("connection refused: {0}")]
ConnectionRefused(String),
#[error("configuration: {0}")]
Configuration(String),
#[error("request build: {0}")]
RequestBuild(String),
#[error("request compression: {0}")]
RequestCompression(String),
#[error("response compression: {0}")]
ResponseCompression(String),
#[error("response decode: {0}")]
ResponseDecode(String),
#[error("session protocol: {0}")]
SessionProtocol(String),
#[error("{0}")]
Other(String),
}
impl ClientError {
pub fn is_retryable(&self) -> bool {
matches!(
self,
Self::Connect(_)
| Self::Timeout
| Self::ConnectionClosedEarly(_)
| Self::RequestCanceled(_)
| Self::UnexpectedEof(_)
| Self::ConnectionReset(_)
| Self::ConnectionAborted(_)
| Self::ConnectionRefused(_)
)
}
pub fn has_no_side_effects(&self) -> bool {
matches!(
self,
Self::Connect(_)
| Self::ConnectionRefused(_)
| Self::Configuration(_)
| Self::RequestBuild(_)
| Self::RequestCompression(_)
)
}
}
impl From<client::HttpError> for ClientError {
fn from(err: client::HttpError) -> Self {
let err_msg = err.to_string();
match err {
client::HttpError::Send(ref send_err) if send_err.is_connect() => {
classify_io_source(&err, &err_msg).unwrap_or(Self::Connect(err_msg))
}
client::HttpError::Send(_) | client::HttpError::Receive(_) => {
classify_hyper_source(&err, &err_msg)
.or_else(|| classify_io_source(&err, &err_msg))
.unwrap_or(Self::Other(err_msg))
}
client::HttpError::RequestBuild(message) => Self::RequestBuild(message),
client::HttpError::RequestCompression(message) => Self::RequestCompression(message),
client::HttpError::ResponseCompression(message) => Self::ResponseCompression(message),
client::HttpError::ResponseDecode(error) => Self::ResponseDecode(error.to_string()),
client::HttpError::Timeout => Self::Timeout,
}
}
}
fn classify_hyper_source(err: &client::HttpError, err_msg: &str) -> Option<ClientError> {
let hyper_err = source_err::<hyper::Error>(err)?;
let err_msg = format!("{hyper_err} -> {err_msg}");
if hyper_err.is_incomplete_message() {
Some(ClientError::ConnectionClosedEarly(err_msg))
} else if hyper_err.is_canceled() {
Some(ClientError::RequestCanceled(err_msg))
} else {
None
}
}
fn classify_io_source(err: &client::HttpError, err_msg: &str) -> Option<ClientError> {
let io_err = source_err::<std::io::Error>(err)?;
let err_msg = format!("{io_err} -> {err_msg}");
Some(match io_err.kind() {
std::io::ErrorKind::UnexpectedEof => ClientError::UnexpectedEof(err_msg),
std::io::ErrorKind::ConnectionReset => ClientError::ConnectionReset(err_msg),
std::io::ErrorKind::ConnectionAborted => ClientError::ConnectionAborted(err_msg),
std::io::ErrorKind::ConnectionRefused => ClientError::ConnectionRefused(err_msg),
_ => return None,
})
}
fn source_err<T: std::error::Error + 'static>(err: &dyn std::error::Error) -> Option<&T> {
let mut source = err.source();
while let Some(err) = source {
if let Some(err) = err.downcast_ref::<T>() {
return Some(err);
}
source = err.source();
}
None
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum AppendConditionFailed {
#[error("fencing token mismatch, expected: {0}")]
FencingTokenMismatch(FencingToken),
#[error("sequence number mismatch, expected: {0}")]
SeqNumMismatch(u64),
}
impl From<api::stream::AppendConditionFailed> for AppendConditionFailed {
fn from(value: api::stream::AppendConditionFailed) -> Self {
match value {
api::stream::AppendConditionFailed::FencingTokenMismatch(token) => {
Self::FencingTokenMismatch(FencingToken::from_server(token.to_string()))
}
api::stream::AppendConditionFailed::SeqNumMismatch(seq) => Self::SeqNumMismatch(seq),
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum RequestError {
#[error(transparent)]
Client(#[from] ClientError),
#[error(transparent)]
Server(#[from] ServerError),
#[error("malformed access token: {0}")]
MalformedAccessToken(String),
#[cfg(feature = "_hidden")]
#[doc(hidden)]
#[error("access token provider failed: {0}")]
AccessTokenProvider(crate::types::AccessTokenProviderError),
#[error(transparent)]
Validation(#[from] ValidationError),
}
impl RequestError {
pub fn is_retryable(&self) -> bool {
match self {
Self::Client(error) => error.is_retryable(),
Self::Server(error) => error.is_retryable(),
#[cfg(feature = "_hidden")]
Self::AccessTokenProvider(error) => error.is_retryable(),
Self::MalformedAccessToken(_) | Self::Validation(_) => false,
}
}
pub fn has_no_side_effects(&self) -> bool {
match self {
Self::Client(error) => error.has_no_side_effects(),
Self::Server(error) => error.has_no_side_effects(),
#[cfg(feature = "_hidden")]
Self::AccessTokenProvider(_) => true,
Self::MalformedAccessToken(_) | Self::Validation(_) => true,
}
}
pub fn server_error(&self) -> Option<&ServerError> {
match self {
Self::Server(error) => Some(error),
_ => None,
}
}
pub(crate) fn is_authentication_error(&self) -> bool {
matches!(
self,
Self::Server(error)
if error.status == StatusCode::UNAUTHORIZED && error.code == "authn"
)
}
}
impl From<ApiError> for RequestError {
fn from(error: ApiError) -> Self {
match error {
ApiError::Client(error) => Self::Client(error),
ApiError::ProtoDecode(error) => {
Self::Client(ClientError::ResponseDecode(error.to_string()))
}
ApiError::TerminalDecode(error) => {
Self::Client(ClientError::SessionProtocol(error.to_string()))
}
ApiError::MalformedAccessToken(error) => Self::MalformedAccessToken(error),
#[cfg(feature = "_hidden")]
ApiError::AccessTokenProvider(error) => Self::AccessTokenProvider(error),
ApiError::Compression(error) => {
Self::Client(ClientError::ResponseCompression(error.to_string()))
}
ApiError::Server(status, response) => {
Self::Server(ServerError::from_api(status, response))
}
other => Self::Client(ClientError::Other(other.to_string())),
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum ReadError {
#[error(transparent)]
Request(#[from] RequestError),
#[error("read from an unwritten position. current tail: {0}")]
ReadUnwritten(StreamPosition),
}
impl ReadError {
pub fn is_retryable(&self) -> bool {
matches!(self, Self::Request(error) if error.is_retryable())
}
pub fn request_error(&self) -> Option<&RequestError> {
match self {
Self::Request(error) => Some(error),
Self::ReadUnwritten(_) => None,
}
}
}
impl From<ApiError> for ReadError {
fn from(error: ApiError) -> Self {
match error {
ApiError::ReadUnwritten(tail) => Self::ReadUnwritten(tail.tail.into()),
other => Self::Request(other.into()),
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum AppendError {
#[error(transparent)]
Request(#[from] RequestError),
#[error(transparent)]
ConditionFailed(#[from] AppendConditionFailed),
}
impl AppendError {
pub fn is_retryable(&self) -> bool {
matches!(self, Self::Request(error) if error.is_retryable())
}
pub fn has_no_side_effects(&self) -> bool {
match self {
Self::Request(error) => error.has_no_side_effects(),
Self::ConditionFailed(_) => true,
}
}
pub fn request_error(&self) -> Option<&RequestError> {
match self {
Self::Request(error) => Some(error),
Self::ConditionFailed(_) => None,
}
}
}
impl From<ApiError> for AppendError {
fn from(error: ApiError) -> Self {
match error {
ApiError::AppendConditionFailed(condition) => Self::ConditionFailed(condition.into()),
other => Self::Request(other.into()),
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum ProducerError {
#[error(transparent)]
Append(#[from] AppendSessionError),
#[error(transparent)]
Validation(#[from] ValidationError),
#[error("producer already closed")]
ProducerClosed,
#[error("producer is closing")]
ProducerClosing,
#[error("producer dropped without calling close")]
ProducerDropped,
}
impl ProducerError {
pub fn is_retryable(&self) -> bool {
match self {
Self::Append(error) => error.is_retryable(),
Self::Validation(_)
| Self::ProducerClosed
| Self::ProducerClosing
| Self::ProducerDropped => false,
}
}
pub fn has_no_side_effects(&self) -> bool {
match self {
Self::Append(error) => error.has_no_side_effects(),
Self::Validation(_) | Self::ProducerClosed | Self::ProducerClosing => true,
Self::ProducerDropped => false,
}
}
pub fn request_error(&self) -> Option<&RequestError> {
match self {
Self::Append(error) => error.request_error(),
Self::Validation(_)
| Self::ProducerClosed
| Self::ProducerClosing
| Self::ProducerDropped => None,
}
}
}
#[derive(Debug, Clone, thiserror::Error)]
#[error("{code}: {message}")]
#[non_exhaustive]
pub struct ServerError {
pub status: StatusCode,
pub code: String,
pub message: String,
}
impl ServerError {
pub(crate) fn from_api(status: StatusCode, response: ServerErrorBody) -> Self {
Self {
status,
code: response.code,
message: response.message,
}
}
pub fn known_code(&self) -> Option<ErrorCode> {
self.code.parse().ok()
}
pub fn is_retryable(&self) -> bool {
server_error_is_retryable(self.status, &self.code)
}
pub fn has_no_side_effects(&self) -> bool {
server_error_has_no_side_effects(self.status, &self.code)
}
}
pub(crate) fn server_error_is_retryable(status: StatusCode, code: &str) -> bool {
match code.parse::<ErrorCode>() {
Ok(code) if code.status() == status => code.is_retryable(),
Ok(_) => false,
Err(_) => matches!(
status,
StatusCode::REQUEST_TIMEOUT
| StatusCode::TOO_MANY_REQUESTS
| StatusCode::INTERNAL_SERVER_ERROR
| StatusCode::BAD_GATEWAY
| StatusCode::SERVICE_UNAVAILABLE
| StatusCode::GATEWAY_TIMEOUT
),
}
}
pub(crate) fn server_error_has_no_side_effects(status: StatusCode, code: &str) -> bool {
code.parse::<ErrorCode>()
.is_ok_and(|code| code.status() == status && code.has_no_side_effects())
}
#[cfg(test)]
mod tests {
use super::*;
fn response(status: StatusCode, code: &str) -> ServerError {
ServerError::from_api(
status,
ServerErrorBody {
code: code.to_owned(),
message: "test".to_owned(),
},
)
}
#[test]
fn error_response_preserves_raw_and_known_codes() {
let known = response(StatusCode::NOT_FOUND, "basin_not_found");
assert_eq!(known.code, "basin_not_found");
assert_eq!(known.message, "test");
assert_eq!(known.known_code(), Some(ErrorCode::BasinNotFound));
assert!(known.to_string().contains("basin_not_found"));
let unknown = response(StatusCode::BAD_REQUEST, "introduced_by_a_newer_server");
assert_eq!(unknown.known_code(), None);
assert_eq!(unknown.code, "introduced_by_a_newer_server");
}
#[test]
fn server_classification_fails_closed_on_status_mismatch() {
let mismatch = response(StatusCode::INTERNAL_SERVER_ERROR, "rate_limited");
assert!(!mismatch.is_retryable());
assert!(!mismatch.has_no_side_effects());
}
#[test]
fn unknown_codes_retain_retryable_status_fallback() {
let unknown = response(StatusCode::SERVICE_UNAVAILABLE, "future_server_error");
assert!(unknown.is_retryable());
assert!(!unknown.has_no_side_effects());
}
#[test]
fn internal_client_errors_preserve_the_failure_stage() {
assert!(matches!(
ClientError::from(client::HttpError::RequestBuild("bad request".to_owned())),
ClientError::RequestBuild(message) if message == "bad request"
));
assert!(matches!(
ClientError::from(client::HttpError::RequestCompression("encode".to_owned())),
ClientError::RequestCompression(message) if message == "encode"
));
assert!(matches!(
ClientError::from(client::HttpError::ResponseCompression("decode".to_owned())),
ClientError::ResponseCompression(message) if message == "decode"
));
let json_error = serde_json::from_slice::<serde_json::Value>(b"{")
.expect_err("invalid JSON should fail");
assert!(matches!(
ClientError::from(client::HttpError::ResponseDecode(json_error)),
ClientError::ResponseDecode(_)
));
}
#[test]
fn nested_errors_expose_request_and_server_errors() {
let append = AppendError::Request(RequestError::Server(response(
StatusCode::CONFLICT,
"transaction_conflict",
)));
assert!(append.is_retryable());
assert!(append.has_no_side_effects());
let request = append.request_error().expect("request error");
assert!(matches!(request, RequestError::Server(_)));
let server = request.server_error().expect("server error");
assert_eq!(server.known_code(), Some(ErrorCode::TransactionConflict));
}
}