use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::Arc;
use tracing::error;
use crate::RequestId;
use crate::builder::SharedHandler;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ConnectionFailureCategory {
Framing,
Protocol,
Transport,
PanicIsolation,
Overload,
Close,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionDirection {
Inbound,
Outbound,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ConnectionRequestId {
Number(i32),
String,
}
impl From<&RequestId> for ConnectionRequestId {
fn from(value: &RequestId) -> Self {
match value {
RequestId::Number(number) => Self::Number(*number),
RequestId::String(_) => Self::String,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct ConnectionFailureContext {
pub connection_id: u64,
pub direction: Option<ConnectionDirection>,
pub method: Option<String>,
pub request_id: Option<ConnectionRequestId>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct ConnectionFailure {
pub category: ConnectionFailureCategory,
pub context: ConnectionFailureContext,
}
pub(crate) type ErrorHook = Arc<dyn SharedHandler<(ConnectionFailure,), ()> + Send + Sync>;
#[derive(Clone)]
pub(crate) struct FailureReporter {
hook: Option<ErrorHook>,
connection_id: u64,
}
impl FailureReporter {
pub(crate) fn new(hook: Option<ErrorHook>, connection_id: u64) -> Self {
Self {
hook,
connection_id,
}
}
pub(crate) fn report(
&self,
category: ConnectionFailureCategory,
direction: Option<ConnectionDirection>,
method: Option<&str>,
request_id: Option<&RequestId>,
) {
let Some(hook) = &self.hook else { return };
let failure = ConnectionFailure {
category,
context: ConnectionFailureContext {
connection_id: self.connection_id,
direction,
method: method.map(str::to_owned),
request_id: request_id.map(ConnectionRequestId::from),
},
};
if catch_unwind(AssertUnwindSafe(|| hook.invoke((failure,)))).is_err() {
error!("panic isolated while invoking connection error hook");
}
}
pub(crate) fn report_unvalidated_inbound_method(
&self,
category: ConnectionFailureCategory,
request_id: Option<&RequestId>,
) {
self.report(
category,
Some(ConnectionDirection::Inbound),
None,
request_id,
);
}
}