use crate::axum::http::StatusCode;
#[derive(Debug)]
pub enum RealtimeError {
Origin,
Unauthorized,
ConnectionLimit,
Protocol {
hint: ProtocolHint,
},
Channel(ChannelError),
Shutdown { remaining: usize },
}
impl std::fmt::Display for RealtimeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Origin => f.write_str("realtime origin not authorized"),
Self::Unauthorized => f.write_str("realtime authorization denied"),
Self::ConnectionLimit => f.write_str("realtime connection limit reached"),
Self::Protocol { hint } => {
f.write_str("realtime protocol error: ")?;
std::fmt::Display::fmt(hint, f)
}
Self::Channel(source) => {
f.write_str("realtime channel error: ")?;
std::fmt::Display::fmt(source, f)
}
Self::Shutdown { remaining } => {
f.write_str("realtime drain timed out: ")?;
std::fmt::Display::fmt(remaining, f)?;
f.write_str(" connections remaining")
}
}
}
}
impl std::error::Error for RealtimeError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Channel(source) => Some(source),
Self::Origin
| Self::Unauthorized
| Self::ConnectionLimit
| Self::Protocol { .. }
| Self::Shutdown { .. } => None,
}
}
}
impl From<ChannelError> for RealtimeError {
fn from(source: ChannelError) -> Self {
Self::Channel(source)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProtocolHint {
Oversize,
Malformed,
Utf8,
Stream,
}
impl ProtocolHint {
const fn as_str(self) -> &'static str {
match self {
Self::Oversize => "oversize",
Self::Malformed => "malformed",
Self::Utf8 => "utf8",
Self::Stream => "stream",
}
}
}
impl std::fmt::Display for ProtocolHint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ChannelError {
Lagged,
Closed,
Full,
}
impl std::fmt::Display for ChannelError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Lagged => f.write_str("channel subscriber lagged"),
Self::Closed => f.write_str("channel closed"),
Self::Full => f.write_str("channel full"),
}
}
}
impl std::error::Error for ChannelError {}
#[must_use]
pub fn admission_status(error: &RealtimeError) -> StatusCode {
match error {
RealtimeError::Origin | RealtimeError::Unauthorized => StatusCode::FORBIDDEN,
RealtimeError::ConnectionLimit => StatusCode::SERVICE_UNAVAILABLE,
RealtimeError::Protocol { .. }
| RealtimeError::Channel(_)
| RealtimeError::Shutdown { .. } => StatusCode::INTERNAL_SERVER_ERROR,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::error::Error as _;
#[test]
fn protocol_hint_display_is_stable_and_safe() {
assert_eq!(ProtocolHint::Oversize.to_string(), "oversize");
assert_eq!(ProtocolHint::Malformed.to_string(), "malformed");
assert_eq!(ProtocolHint::Utf8.to_string(), "utf8");
assert_eq!(ProtocolHint::Stream.to_string(), "stream");
}
#[test]
fn admission_status_maps_origin_and_authz_to_403() {
assert_eq!(
admission_status(&RealtimeError::Origin),
StatusCode::FORBIDDEN
);
assert_eq!(
admission_status(&RealtimeError::Unauthorized),
StatusCode::FORBIDDEN
);
}
#[test]
fn admission_status_maps_connection_limit_to_503() {
assert_eq!(
admission_status(&RealtimeError::ConnectionLimit),
StatusCode::SERVICE_UNAVAILABLE
);
}
#[test]
fn admission_status_maps_post_upgrade_errors_to_500() {
assert_eq!(
admission_status(&RealtimeError::Protocol {
hint: ProtocolHint::Oversize
}),
StatusCode::INTERNAL_SERVER_ERROR
);
assert_eq!(
admission_status(&RealtimeError::Channel(ChannelError::Lagged)),
StatusCode::INTERNAL_SERVER_ERROR
);
assert_eq!(
admission_status(&RealtimeError::Shutdown { remaining: 1 }),
StatusCode::INTERNAL_SERVER_ERROR
);
}
#[test]
fn display_does_not_carry_attacker_bytes() {
let err = RealtimeError::Protocol {
hint: ProtocolHint::Malformed,
};
assert_eq!(err.to_string(), "realtime protocol error: malformed");
}
#[test]
fn channel_error_source_is_itself() {
let err = RealtimeError::from(ChannelError::Closed);
assert!(err.source().is_some());
assert_eq!(err.source().unwrap().to_string(), "channel closed");
}
}