#[cfg(feature = "websocket")]
use std::net::SocketAddr;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use futures_util::FutureExt;
use tracing::error;
#[cfg(feature = "websocket")]
use tracing::info;
use crate::ack::{AckManager, SharedAckManager};
use crate::broker::{Broker, InMemoryBroker};
use crate::config::ServerConfig;
use crate::connection::Connection;
use crate::error::ConfigError;
use crate::error::Result;
use crate::error::RiftError;
use crate::metrics::Metrics;
use crate::session::AuthProvider;
use crate::session::resume::ResumeManager;
use crate::session::store::SessionStore;
use crate::transport::TransportConnection;
#[cfg(feature = "websocket")]
use crate::transport::{Transport, TransportListener};
#[cfg(feature = "websocket")]
use crate::transport::websocket::WebSocketTransport;
#[cfg(feature = "websocket")]
type ListenerFuture =
std::pin::Pin<Box<dyn std::future::Future<Output = Result<Box<dyn TransportListener>>> + Send>>;
#[cfg(feature = "websocket")]
trait TransportFactory: Send + Sync {
fn build(&self, addr: SocketAddr) -> ListenerFuture;
}
#[cfg(feature = "websocket")]
struct WebSocketFactory {
max_message_size: usize,
}
#[cfg(feature = "websocket")]
impl TransportFactory for WebSocketFactory {
fn build(&self, addr: SocketAddr) -> ListenerFuture {
let transport = WebSocketTransport::new().with_max_message_size(self.max_message_size);
Box::pin(async move { transport.bind(addr).await })
}
}
pub struct RiftServerBuilder {
config: ServerConfig,
auth: Option<Arc<dyn AuthProvider>>,
broker: Option<Arc<dyn Broker>>,
#[cfg(feature = "websocket")]
transport: Option<Arc<dyn TransportFactory>>,
metrics: Option<Arc<Metrics>>,
session_store: Option<SessionStore>,
resume_manager: Option<Arc<ResumeManager>>,
}
impl RiftServerBuilder {
pub fn new() -> Self {
Self {
config: ServerConfig::default(),
auth: None,
broker: None,
#[cfg(feature = "websocket")]
transport: None,
metrics: None,
session_store: None,
resume_manager: None,
}
}
pub fn config(mut self, config: ServerConfig) -> Self {
self.config = config;
self
}
pub fn auth(mut self, auth: Arc<dyn AuthProvider>) -> Self {
self.auth = Some(auth);
self
}
pub fn broker(mut self, broker: Arc<dyn Broker>) -> Self {
self.broker = Some(broker);
self
}
#[cfg(feature = "websocket")]
pub fn websocket_transport(mut self) -> Self {
self.transport = Some(Arc::new(WebSocketFactory {
max_message_size: 0,
}));
self
}
pub fn metrics(mut self, metrics: Arc<Metrics>) -> Self {
self.metrics = Some(metrics);
self
}
pub fn session_store(mut self, store: SessionStore) -> Self {
self.session_store = Some(store);
self
}
pub fn resume_manager(mut self, rm: Arc<ResumeManager>) -> Self {
self.resume_manager = Some(rm);
self
}
#[cfg(feature = "redis")]
pub fn redis_broker(mut self, broker: Arc<dyn Broker>) -> Self {
self.broker = Some(broker);
self
}
pub fn build(self) -> Result<RiftServer> {
let metrics = self.metrics.unwrap_or_else(|| Arc::new(Metrics::new()));
let config_max_payload = self.config.max_payload_bytes;
let broker = self.broker.unwrap_or_else(|| {
let topic_profile: crate::topic::TopicProfile =
self.config.default_topic_profile.clone();
Arc::new(InMemoryBroker::new(
topic_profile,
self.config.dedupe_window,
config_max_payload,
))
});
let auth = self
.auth
.unwrap_or_else(|| Arc::new(crate::session::TokenAuth::new()));
#[cfg(feature = "websocket")]
let transport = self.transport.as_ref().map(|_| {
Arc::new(WebSocketFactory {
max_message_size: self.config.max_payload_bytes,
}) as Arc<dyn TransportFactory>
});
let session_store = self.session_store.unwrap_or_default();
let resume_manager = self
.resume_manager
.unwrap_or_else(|| Arc::new(ResumeManager::new()));
let ack_manager: SharedAckManager = Arc::new(AckManager::new());
let gc_shutdown = Arc::new(tokio::sync::Notify::new());
let gc_broker = broker.clone();
let gc_session_store = session_store.clone();
let gc_ack = ack_manager.clone();
let gc_idle_timeout = self.config.idle_timeout;
let gc_notify = gc_shutdown.clone();
tokio::spawn(async move {
run_maintenance(
gc_broker,
gc_session_store,
gc_ack,
gc_idle_timeout,
gc_notify,
)
.await;
});
Ok(RiftServer {
config: self.config,
auth,
broker,
metrics,
#[cfg(feature = "websocket")]
transport,
next_conn_id: Arc::new(AtomicU64::new(1)),
session_store,
resume_manager,
ack_manager,
gc_shutdown,
})
}
}
impl Default for RiftServerBuilder {
fn default() -> Self {
Self::new()
}
}
pub struct RiftServer {
pub config: ServerConfig,
auth: Arc<dyn AuthProvider>,
broker: Arc<dyn Broker>,
metrics: Arc<Metrics>,
#[cfg(feature = "websocket")]
transport: Option<Arc<dyn TransportFactory>>,
next_conn_id: Arc<AtomicU64>,
session_store: SessionStore,
resume_manager: Arc<ResumeManager>,
ack_manager: SharedAckManager,
gc_shutdown: Arc<tokio::sync::Notify>,
}
impl RiftServer {
pub fn builder() -> RiftServerBuilder {
RiftServerBuilder::new()
}
#[cfg(feature = "websocket")]
pub async fn run(&self, addr: SocketAddr, shutdown: Arc<tokio::sync::Notify>) -> Result<()> {
let transport = self.transport.as_ref().ok_or_else(|| {
RiftError::Config(ConfigError::Invalid {
field: "transport",
message:
"no transport configured; call builder.websocket_transport() before build(), \
or use accept_and_spawn() for framework mode"
.to_string(),
})
})?;
let mut listener = transport.build(addr).await?;
info!(addr = ?listener.local_addr()?, "rift server listening");
loop {
tokio::select! {
_ = shutdown.notified() => {
info!("shutdown signaled");
self.gc_shutdown.notify_waiters();
return Ok(());
}
res = listener.accept() => {
match res {
Ok(conn) => {
self.spawn_connection(conn);
}
Err(e) => {
error!("accept error: {}", e);
}
}
}
}
}
}
pub fn accept_and_spawn(&self, transport: Box<dyn TransportConnection>) {
self.spawn_connection(transport);
}
pub fn shutdown(&self) {
self.gc_shutdown.notify_waiters();
}
fn spawn_connection(&self, mut transport: Box<dyn TransportConnection>) {
let max = self.config.max_connections;
if max > 0 {
let current = self
.metrics
.active_connections
.load(std::sync::atomic::Ordering::SeqCst);
if current as usize >= max {
tracing::warn!(max, "connection limit reached, rejecting new connection");
tokio::spawn(async move {
let _ = transport
.close(
crate::protocol::close::CloseCode::ServerOverloaded,
"server at connection limit",
)
.await;
});
return;
}
}
let id = self
.next_conn_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let ack_manager = self.ack_manager.clone();
let offset_tracker = self.session_store.offset_tracker().clone();
let connection = Connection::new(
id,
self.broker.clone(),
self.auth.clone(),
self.config.clone(),
self.metrics.clone(),
ack_manager,
self.resume_manager.clone(),
offset_tracker,
self.session_store.clone(),
);
tokio::spawn(async move {
let result = AssertUnwindSafe(connection.run(transport))
.catch_unwind()
.await;
match result {
Ok(Ok(())) => {
tracing::debug!(conn = id, "connection ended cleanly");
}
Ok(Err(RiftError::Session(crate::error::SessionReject::IdleTimeout))) => {
tracing::debug!(conn = id, "connection closed due to idle timeout");
}
Ok(Err(e)) => {
error!(conn = id, "connection ended with error: {}", e);
}
Err(panic) => {
error!(conn = id, "connection task panicked: {:?}", panic);
}
}
});
}
}
const MAINTENANCE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(30);
async fn run_maintenance(
broker: Arc<dyn Broker>,
session_store: SessionStore,
ack_manager: SharedAckManager,
idle_timeout: std::time::Duration,
shutdown: Arc<tokio::sync::Notify>,
) {
let mut interval = tokio::time::interval(MAINTENANCE_INTERVAL);
interval.tick().await;
loop {
tokio::select! {
_ = shutdown.notified() => {
tracing::info!("maintenance task shutting down");
break;
}
_ = interval.tick() => {}
}
let swept = broker.maintain().await;
let sessions_expired = session_store.expire_sessions(idle_timeout);
let acks_reaped = ack_manager.reap_all_timeouts();
if swept > 0 || sessions_expired > 0 || acks_reaped > 0 {
tracing::debug!(swept, sessions_expired, acks_reaped, "maintenance sweep");
}
}
}