use std::fmt::Debug;
use std::sync::Arc;
use std::time::Duration;
use crate::async_trait;
use crate::conn::{ConnCtrl, SocketAddr};
#[derive(Default, Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TransProto {
#[default]
Tcp,
Quic,
}
#[derive(Clone, Debug)]
pub struct FuseInfo {
pub trans_proto: TransProto,
pub remote_addr: SocketAddr,
pub local_addr: SocketAddr,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct FuseConfig {
pub tls_handshake_timeout: Option<Duration>,
pub http1_header_timeout: Option<Duration>,
pub connection_idle_timeout: Option<Duration>,
pub write_stall_timeout: Option<Duration>,
pub request_body_timeout: Option<Duration>,
}
impl Default for FuseConfig {
fn default() -> Self {
Self {
tls_handshake_timeout: Some(Duration::from_secs(10)),
http1_header_timeout: Some(Duration::from_secs(30)),
connection_idle_timeout: None,
write_stall_timeout: None,
request_body_timeout: None,
}
}
}
impl FuseConfig {
#[must_use]
pub const fn strict() -> Self {
Self {
tls_handshake_timeout: Some(Duration::from_secs(10)),
http1_header_timeout: Some(Duration::from_secs(30)),
connection_idle_timeout: Some(Duration::from_secs(30)),
write_stall_timeout: Some(Duration::from_secs(30)),
request_body_timeout: Some(Duration::from_secs(60)),
}
}
#[must_use]
pub const fn disabled() -> Self {
Self {
tls_handshake_timeout: None,
http1_header_timeout: None,
connection_idle_timeout: None,
write_stall_timeout: None,
request_body_timeout: None,
}
}
#[must_use]
pub fn with_tls_handshake_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
self.tls_handshake_timeout = timeout.into();
self
}
#[must_use]
pub fn with_http1_header_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
self.http1_header_timeout = timeout.into();
self
}
#[must_use]
pub fn with_connection_idle_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
self.connection_idle_timeout = timeout.into();
self
}
#[must_use]
pub fn with_write_stall_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
self.write_stall_timeout = timeout.into();
self
}
#[must_use]
pub fn with_request_body_timeout(mut self, timeout: impl Into<Option<Duration>>) -> Self {
self.request_body_timeout = timeout.into();
self
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FuseAction {
Accept(FuseConfig),
Reject,
}
pub trait ConnObserver: Send + Sync + 'static {
fn on_read(&self, bytes: usize) {
let _ = bytes;
}
fn on_write(&self, bytes: usize) {
let _ = bytes;
}
}
pub type ArcConnObserver = Arc<dyn ConnObserver>;
#[async_trait]
pub trait FusePolicy: Send + Sync + 'static {
async fn decide(&self, info: &FuseInfo) -> FuseAction;
fn observe(&self, info: &FuseInfo, ctrl: &ConnCtrl) -> Option<ArcConnObserver> {
let _ = (info, ctrl);
None
}
}
#[async_trait]
impl FusePolicy for FuseConfig {
async fn decide(&self, _info: &FuseInfo) -> FuseAction {
FuseAction::Accept(*self)
}
}
#[async_trait]
impl<F> FusePolicy for F
where
F: Fn(&FuseInfo) -> FuseAction + Send + Sync + 'static,
{
async fn decide(&self, info: &FuseInfo) -> FuseAction {
self(info)
}
}
pub type ArcFusePolicy = Arc<dyn FusePolicy>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_enables_only_the_harmless_timeouts() {
let config = FuseConfig::default();
assert!(config.tls_handshake_timeout.is_some());
assert!(config.http1_header_timeout.is_some());
assert!(config.connection_idle_timeout.is_none());
assert!(config.write_stall_timeout.is_none());
assert!(config.request_body_timeout.is_none());
}
#[test]
fn strict_enables_every_timeout() {
let config = FuseConfig::strict();
assert!(config.tls_handshake_timeout.is_some());
assert!(config.http1_header_timeout.is_some());
assert!(config.connection_idle_timeout.is_some());
assert!(config.write_stall_timeout.is_some());
assert!(config.request_body_timeout.is_some());
}
#[tokio::test]
async fn async_policy_can_await_before_admission() {
struct Blocklist;
#[async_trait]
impl FusePolicy for Blocklist {
async fn decide(&self, info: &FuseInfo) -> FuseAction {
tokio::task::yield_now().await;
if info.remote_addr.as_ipv4().is_some() {
FuseAction::Reject
} else {
FuseAction::Accept(FuseConfig::disabled())
}
}
}
let policy: ArcFusePolicy = Arc::new(Blocklist);
let addr = std::net::SocketAddr::from(([127, 0, 0, 1], 0));
let info = FuseInfo {
trans_proto: TransProto::Tcp,
remote_addr: addr.into(),
local_addr: addr.into(),
};
assert_eq!(policy.decide(&info).await, FuseAction::Reject);
}
}