use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TerminalScope {
Connection,
Request,
}
pub const WIRE_ERROR_CODES: &[&str] = &[
"unsupported_version",
"identity_rejected",
"malformed_frame",
"frame_too_large",
"subscriber_overflow",
"subscription_revoked",
"context_rejected",
"peer_class_denied",
"subscription_denied",
"already_subscribed",
"cursor_expired",
"in_flight_limit_exceeded",
"deadline_exceeded",
"cancelled",
"shutting_down",
"internal",
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum WireErrorCode {
UnsupportedVersion,
IdentityRejected,
MalformedFrame,
FrameTooLarge,
SubscriberOverflow,
SubscriptionRevoked,
ContextRejected,
PeerClassDenied,
SubscriptionDenied,
AlreadySubscribed,
CursorExpired,
InFlightLimitExceeded,
DeadlineExceeded,
Cancelled,
ShuttingDown,
#[serde(other)]
Internal,
}
impl WireErrorCode {
pub const fn terminal_scope(self) -> TerminalScope {
use WireErrorCode::*;
match self {
UnsupportedVersion | IdentityRejected | MalformedFrame | FrameTooLarge
| SubscriberOverflow | SubscriptionRevoked => TerminalScope::Connection,
ContextRejected
| PeerClassDenied
| SubscriptionDenied
| AlreadySubscribed
| CursorExpired
| InFlightLimitExceeded
| DeadlineExceeded
| Cancelled
| ShuttingDown
| Internal => TerminalScope::Request,
}
}
pub const fn as_str(self) -> &'static str {
use WireErrorCode::*;
match self {
UnsupportedVersion => "unsupported_version",
IdentityRejected => "identity_rejected",
MalformedFrame => "malformed_frame",
FrameTooLarge => "frame_too_large",
SubscriberOverflow => "subscriber_overflow",
SubscriptionRevoked => "subscription_revoked",
ContextRejected => "context_rejected",
PeerClassDenied => "peer_class_denied",
SubscriptionDenied => "subscription_denied",
AlreadySubscribed => "already_subscribed",
CursorExpired => "cursor_expired",
InFlightLimitExceeded => "in_flight_limit_exceeded",
DeadlineExceeded => "deadline_exceeded",
Cancelled => "cancelled",
ShuttingDown => "shutting_down",
Internal => "internal",
}
}
}
impl std::fmt::Display for WireErrorCode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn unrecognized_code_falls_back_to_internal() {
let code: WireErrorCode = serde_json::from_str(r#""bogus""#).unwrap();
assert_eq!(code, WireErrorCode::Internal);
assert_eq!(code.terminal_scope(), TerminalScope::Request);
}
#[test]
fn unrecognized_code_inside_an_error_frame_falls_back_to_internal() {
let payload = br#"{"kind":"error","code":"bogus","message":"future code"}"#;
let frame = crate::codec::decode_payload(payload).unwrap();
match frame {
crate::frame::Frame::Error {
id,
code,
unrecognized_code,
..
} => {
assert_eq!(code, WireErrorCode::Internal);
assert!(id.is_none());
assert_eq!(unrecognized_code.as_deref(), Some("bogus"));
}
other => panic!("expected an error frame, got {other:?}"),
}
}
#[test]
fn wire_error_codes_matches_the_code_enum() {
for entry in WIRE_ERROR_CODES {
let json = format!("\"{entry}\"");
let parsed: WireErrorCode = serde_json::from_str(&json).unwrap();
assert_eq!(
parsed.as_str(),
*entry,
"WIRE_ERROR_CODES entry {entry:?} names no variant \
(deserialized to {parsed:?})"
);
}
let mut deduped: Vec<&str> = WIRE_ERROR_CODES.to_vec();
deduped.sort_unstable();
deduped.dedup();
assert_eq!(
deduped.len(),
WIRE_ERROR_CODES.len(),
"WIRE_ERROR_CODES contains duplicates"
);
use WireErrorCode::*;
let variants = [
UnsupportedVersion,
IdentityRejected,
MalformedFrame,
FrameTooLarge,
SubscriberOverflow,
SubscriptionRevoked,
ContextRejected,
PeerClassDenied,
SubscriptionDenied,
AlreadySubscribed,
CursorExpired,
InFlightLimitExceeded,
DeadlineExceeded,
Cancelled,
ShuttingDown,
Internal,
];
for code in variants {
match code {
UnsupportedVersion
| IdentityRejected
| MalformedFrame
| FrameTooLarge
| SubscriberOverflow
| SubscriptionRevoked
| ContextRejected
| PeerClassDenied
| SubscriptionDenied
| AlreadySubscribed
| CursorExpired
| InFlightLimitExceeded
| DeadlineExceeded
| Cancelled
| ShuttingDown
| Internal => {}
}
assert!(
WIRE_ERROR_CODES.contains(&code.as_str()),
"WIRE_ERROR_CODES is missing {code:?} ({}); its scope pairing \
would silently go unenforced at decode",
code.as_str()
);
}
assert_eq!(
WIRE_ERROR_CODES.len(),
variants.len(),
"WIRE_ERROR_CODES and the variant list disagree on count"
);
}
}