use std::net::SocketAddr;
use tracing::Instrument;
use crate::decode::DecodeLevel;
use crate::server::task::ServerSetting;
use crate::tcp::server::{ServerTask as TcpServerTask, TcpServerConnectionHandler};
mod address_filter;
pub(crate) mod handler;
pub(crate) mod request;
pub(crate) mod response;
pub(crate) mod task;
pub(crate) mod types;
pub(crate) const SERVER_SETTING_CHANNEL_CAPACITY: usize = 8;
use crate::error::Shutdown;
pub use address_filter::*;
pub use handler::*;
pub use types::*;
#[cfg(feature = "enable-tls")]
pub use crate::tcp::tls::server::TlsServerConfig;
#[cfg(feature = "enable-tls")]
pub use crate::tcp::tls::*;
#[derive(Debug)]
pub struct ServerHandle {
tx: tokio::sync::mpsc::Sender<ServerSetting>,
}
pub struct ServerTask<T: RequestHandler> {
inner: ServerTaskInner<T>,
}
enum ServerTaskInner<T: RequestHandler> {
Tcp(
Box<TcpServerTask<T>>,
tokio::sync::mpsc::Receiver<ServerSetting>,
),
#[cfg(feature = "serial")]
Rtu(Box<crate::serial::server::RtuServerTask<T>>),
}
impl<T: RequestHandler> ServerTask<T> {
fn tcp(task: TcpServerTask<T>, commands: tokio::sync::mpsc::Receiver<ServerSetting>) -> Self {
Self {
inner: ServerTaskInner::Tcp(Box::new(task), commands),
}
}
#[cfg(feature = "serial")]
fn rtu(task: crate::serial::server::RtuServerTask<T>) -> Self {
Self {
inner: ServerTaskInner::Rtu(Box::new(task)),
}
}
pub async fn run(self) {
match self.inner {
ServerTaskInner::Tcp(mut task, commands) => task.run(commands).await,
#[cfg(feature = "serial")]
ServerTaskInner::Rtu(mut task) => {
task.run().await;
}
}
}
}
impl ServerHandle {
pub fn new(tx: tokio::sync::mpsc::Sender<ServerSetting>) -> Self {
ServerHandle { tx }
}
pub async fn set_decode_level(&mut self, level: DecodeLevel) -> Result<(), Shutdown> {
self.tx.send(ServerSetting::ChangeDecoding(level)).await?;
Ok(())
}
}
pub async fn spawn_tcp_server_task<T: RequestHandler>(
max_sessions: usize,
addr: SocketAddr,
handlers: ServerHandlerMap<T>,
filter: AddressFilter,
decode: DecodeLevel,
) -> Result<ServerHandle, std::io::Error> {
let listener = tokio::net::TcpListener::bind(addr).await?;
let (handle, task) = create_tcp_server_task(max_sessions, listener, handlers, filter, decode);
tokio::spawn(
task.run()
.instrument(tracing::info_span!("Modbus-Server-TCP", "listen" = ?addr)),
);
Ok(handle)
}
pub fn create_tcp_server_task<T: RequestHandler>(
max_sessions: usize,
listener: tokio::net::TcpListener,
handlers: ServerHandlerMap<T>,
filter: AddressFilter,
decode: DecodeLevel,
) -> (ServerHandle, ServerTask<T>) {
let (tx, rx) = tokio::sync::mpsc::channel(SERVER_SETTING_CHANNEL_CAPACITY);
let task = TcpServerTask::new(
max_sessions,
listener,
handlers,
TcpServerConnectionHandler::Tcp,
filter,
decode,
);
(ServerHandle::new(tx), ServerTask::tcp(task, rx))
}
#[cfg(feature = "serial")]
pub fn spawn_rtu_server_task<T: RequestHandler>(
path: &str,
settings: crate::serial::SerialSettings,
retry: Box<dyn crate::retry::RetryStrategy>,
handlers: ServerHandlerMap<T>,
decode: DecodeLevel,
) -> Result<ServerHandle, std::io::Error> {
let (handle, task) = create_rtu_server_task(path, settings, retry, handlers, decode);
let path = path.to_string();
tokio::spawn(
task.run()
.instrument(tracing::info_span!("Modbus-Server-RTU", "port" = ?path)),
);
Ok(handle)
}
#[cfg(feature = "serial")]
pub fn create_rtu_server_task<T: RequestHandler>(
path: &str,
settings: crate::serial::SerialSettings,
retry: Box<dyn crate::retry::RetryStrategy>,
handlers: ServerHandlerMap<T>,
decode: DecodeLevel,
) -> (ServerHandle, ServerTask<T>) {
let (tx, rx) = tokio::sync::mpsc::channel(SERVER_SETTING_CHANNEL_CAPACITY);
let session = crate::server::task::SessionTask::new(
handlers,
crate::server::task::AuthorizationType::None,
crate::common::frame::FrameWriter::rtu(),
crate::common::frame::FramedReader::rtu_request(),
rx,
decode,
);
let rtu = crate::serial::server::RtuServerTask {
port: path.to_string(),
retry,
settings,
session,
};
(ServerHandle::new(tx), ServerTask::rtu(rtu))
}
#[cfg(feature = "enable-tls")]
pub async fn spawn_tls_server_task<T: RequestHandler>(
max_sessions: usize,
addr: SocketAddr,
handlers: ServerHandlerMap<T>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> Result<ServerHandle, std::io::Error> {
spawn_tls_server_task_impl(
max_sessions,
addr,
handlers,
None,
tls_config,
filter,
decode,
)
.await
}
#[cfg(feature = "enable-tls")]
pub fn create_tls_server_task<T: RequestHandler>(
max_sessions: usize,
listener: tokio::net::TcpListener,
handlers: ServerHandlerMap<T>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> (ServerHandle, ServerTask<T>) {
create_tls_server_task_impl(
max_sessions,
listener,
handlers,
None,
tls_config,
filter,
decode,
)
}
#[cfg(feature = "enable-tls")]
pub async fn spawn_tls_server_task_with_authz<T: RequestHandler>(
max_sessions: usize,
addr: SocketAddr,
handlers: ServerHandlerMap<T>,
auth_handler: std::sync::Arc<dyn AuthorizationHandler>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> Result<ServerHandle, std::io::Error> {
spawn_tls_server_task_impl(
max_sessions,
addr,
handlers,
Some(auth_handler),
tls_config,
filter,
decode,
)
.await
}
#[cfg(feature = "enable-tls")]
pub fn create_tls_server_task_with_authz<T: RequestHandler>(
max_sessions: usize,
listener: tokio::net::TcpListener,
handlers: ServerHandlerMap<T>,
auth_handler: std::sync::Arc<dyn AuthorizationHandler>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> (ServerHandle, ServerTask<T>) {
create_tls_server_task_impl(
max_sessions,
listener,
handlers,
Some(auth_handler),
tls_config,
filter,
decode,
)
}
#[cfg(feature = "enable-tls")]
async fn spawn_tls_server_task_impl<T: RequestHandler>(
max_sessions: usize,
addr: SocketAddr,
handlers: ServerHandlerMap<T>,
auth_handler: Option<std::sync::Arc<dyn AuthorizationHandler>>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> Result<ServerHandle, std::io::Error> {
let listener = tokio::net::TcpListener::bind(addr).await?;
let (handle, task) = create_tls_server_task_impl(
max_sessions,
listener,
handlers,
auth_handler,
tls_config,
filter,
decode,
);
tokio::spawn(
task.run()
.instrument(tracing::info_span!("Modbus-Server-TLS", "listen" = ?addr)),
);
Ok(handle)
}
#[cfg(feature = "enable-tls")]
fn create_tls_server_task_impl<T: RequestHandler>(
max_sessions: usize,
listener: tokio::net::TcpListener,
handlers: ServerHandlerMap<T>,
auth_handler: Option<std::sync::Arc<dyn AuthorizationHandler>>,
tls_config: TlsServerConfig,
filter: AddressFilter,
decode: DecodeLevel,
) -> (ServerHandle, ServerTask<T>) {
let (tx, rx) = tokio::sync::mpsc::channel(SERVER_SETTING_CHANNEL_CAPACITY);
let task = TcpServerTask::new(
max_sessions,
listener,
handlers,
TcpServerConnectionHandler::Tls(tls_config, auth_handler),
filter,
decode,
);
(ServerHandle::new(tx), ServerTask::tcp(task, rx))
}