#[cfg(feature = "websocket")]
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use tracing::error;
#[cfg(feature = "websocket")]
use tracing::info;
use crate::ack::{AckManager, SharedAckManager};
use crate::broker::{Broker, InMemoryBroker};
use crate::config::{DefaultTopicProfile, ServerConfig};
use crate::connection::Connection;
use crate::error::Result;
#[cfg(feature = "websocket")]
use crate::error::{ConfigError, RiftError};
use crate::metrics::Metrics;
use crate::session::AuthProvider;
use crate::session::resume::ResumeManager;
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<Box<dyn TransportFactory>>,
metrics: Option<Arc<Metrics>>,
}
impl RiftServerBuilder {
pub fn new() -> Self {
Self {
config: ServerConfig::default(),
auth: None,
broker: None,
#[cfg(feature = "websocket")]
transport: None,
metrics: 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 {
let max_msg = self.config.max_payload_bytes;
self.transport = Some(Box::new(WebSocketFactory {
max_message_size: max_msg,
}));
self
}
pub fn metrics(mut self, metrics: Arc<Metrics>) -> Self {
self.metrics = Some(metrics);
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().into();
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()));
Ok(RiftServer {
config: self.config,
auth,
broker,
metrics,
#[cfg(feature = "websocket")]
transport: Arc::from(self.transport.ok_or_else(|| {
RiftError::Config(ConfigError::Invalid {
field: "transport",
message: "transport is required for standalone mode".to_string(),
})
})?),
next_conn_id: Arc::new(AtomicU64::new(1)),
})
}
}
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: Arc<dyn TransportFactory>,
next_conn_id: Arc<AtomicU64>,
}
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 mut listener = self.transport.build(addr).await?;
info!(addr = ?listener.local_addr()?, "rift server listening");
loop {
tokio::select! {
_ = shutdown.notified() => {
info!("shutdown signaled");
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);
}
fn spawn_connection(&self, transport: Box<dyn TransportConnection>) {
let id = self
.next_conn_id
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let ack_manager: SharedAckManager = Arc::new(AckManager::new());
let resume = Arc::new(ResumeManager::new());
let connection = Connection::new(
id,
self.broker.clone(),
self.auth.clone(),
self.config.clone(),
self.metrics.clone(),
ack_manager,
resume,
);
tokio::spawn(async move {
if let Err(e) = connection.run(transport).await {
error!(conn = id, "connection ended with error: {}", e);
}
});
}
}
impl From<DefaultTopicProfile> for crate::topic::TopicProfile {
fn from(d: DefaultTopicProfile) -> Self {
Self {
name: "default".into(),
retention: d.retention,
ordering: d.ordering,
max_subscribers: d.max_subscribers,
max_publishers: d.max_publishers,
rate_limit_per_publisher: None,
rate_limit_total: None,
replay_enabled: d.replay_enabled,
snapshot_enabled: d.snapshot_enabled,
replay_window: std::time::Duration::from_secs(300),
}
}
}