use std::time::{Duration, Instant};
use bytes::{BufMut, Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use crate::error::ModbusError;
use crate::error::*;
use crate::frame::{self, Request, Response};
const DRAIN_QUIET_MS: u64 = 10;
const DRAIN_MAX_MS: u64 = 200;
pub(crate) async fn drain_stale_data<S>(
stream: &mut S,
scratch: &mut [u8],
) -> Result<(), ModbusError>
where
S: AsyncRead + Unpin,
{
const QUICK_POLL_MS: u64 = 1;
let quick = Duration::from_millis(QUICK_POLL_MS);
match tokio::time::timeout(quick, stream.read(scratch)).await {
Ok(Ok(0)) => return Err(ModbusError::connection("connection closed")),
Ok(Ok(_)) => {
}
Ok(Err(e)) => return Err(ModbusError::from(e)),
Err(_) => {
return Ok(());
}
}
let per_read = Duration::from_millis(DRAIN_QUIET_MS);
let deadline = Instant::now() + Duration::from_millis(DRAIN_MAX_MS);
loop {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
return Err(ModbusError::connection(
"bus busy: data received during quiet period",
));
}
let wait = remaining.min(per_read);
match tokio::time::timeout(wait, stream.read(scratch)).await {
Ok(Ok(0)) => return Err(ModbusError::connection("connection closed")),
Ok(Ok(_)) => {
}
Ok(Err(e)) => return Err(ModbusError::from(e)),
Err(_) => {
return Ok(());
}
}
}
}
pub(crate) async fn send_frame<S, E>(
stream: &mut S,
write_buf: &mut BytesMut,
slave_id: u8,
timeout: Duration,
request: &Request<'_>,
mut encode_fn: E,
) -> Result<(), ModbusError>
where
S: AsyncWrite + Unpin,
E: FnMut(&[u8], &mut BytesMut),
{
write_buf.clear();
write_buf.put_u8(slave_id);
frame::encode_request_into(request, write_buf)
.map_err(|e| ModbusError::protocol(format!("{PDU_ENCODE_ERROR} {e}")))?;
let send_buf = write_buf.split();
encode_fn(&send_buf, write_buf);
let frame_bytes = write_buf.split().freeze();
tokio::time::timeout(timeout, stream.write_all(&frame_bytes))
.await
.map_err(|_| ModbusError::timeout(SEND_TIMEOUT))?
.map_err(ModbusError::from)?;
Ok(())
}
pub(crate) async fn read_at_least<S>(
stream: &mut S,
read_buf: &mut BytesMut,
deadline: Instant,
min: usize,
) -> Result<usize, ModbusError>
where
S: AsyncRead + Unpin,
{
while read_buf.len() < min {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
break;
}
match tokio::time::timeout(remaining, stream.read_buf(read_buf)).await {
Ok(Ok(0)) => {
if read_buf.is_empty() {
return Err(ModbusError::connection(CONN_CLOSED));
}
break; }
Ok(Ok(_)) => {} Ok(Err(e)) => return Err(ModbusError::from(e)),
Err(_) => {
if !read_buf.is_empty() {
break;
}
return Err(ModbusError::timeout(RECV_TIMEOUT));
}
}
}
Ok(read_buf.len())
}
use crate::reconnect::ReconnectConfig;
use std::future::Future;
pub(crate) async fn run_with_reconnect<Fut, S, R, RFut, E>(
reconnect: Option<&ReconnectConfig>,
mut send: S,
mut rebuild: R,
exceeded_err: E,
) -> Result<Response, ModbusError>
where
Fut: Future<Output = Result<Response, ModbusError>>,
S: FnMut() -> Fut,
R: FnMut() -> RFut,
RFut: Future<Output = bool>,
E: FnOnce(String) -> ModbusError,
{
let Some(cfg) = reconnect else {
return send().await;
};
let mut failures = 0u32;
let mut rebuild_failures = 0u32;
loop {
match send().await {
Ok(rsp) => return Ok(rsp),
Err(e) => {
if !crate::reconnect::is_retryable(&e) {
return Err(e);
}
failures += 1;
if cfg.max_retries() > 0 && failures > cfg.max_retries() {
return Err(exceeded_err(format!(
"reconnect max_retries({}) exceeded: {}",
cfg.max_retries(),
e.detail()
)));
}
tokio::time::sleep(cfg.interval()).await;
if !rebuild().await {
rebuild_failures += 1;
if rebuild_failures >= 3 {
return Err(exceeded_err(
"reconnect: transport rebuild failed 3 consecutive times".into(),
));
}
continue;
}
rebuild_failures = 0;
}
}
}
}
pub(crate) async fn process_server_request<S>(
raw: &Bytes,
slave_id: u8,
service: &S,
rsp_buf: &mut BytesMut,
) -> Option<Bytes>
where
S: crate::server::Service + Send + Sync,
{
if slave_id == 0 {
return None;
}
if raw.is_empty() {
return None;
}
let request = match Request::try_from(raw.clone()) {
Ok(req) => req,
Err(e) => {
log::warn!("server: invalid request from slave {}: {}", slave_id, e);
return None;
}
};
let req_fc = request.function_code().value();
let response = crate::server::context::SLAVE_ID
.scope(slave_id, async { service.call(request).await })
.await;
let response = match response {
Ok(rsp) => rsp,
Err(ex) => {
log::warn!("server: exception from slave {}: {:?}", slave_id, ex);
Response::Exception(req_fc, ex)
}
};
rsp_buf.clear();
rsp_buf.put_u8(slave_id);
if frame::encode_response_into(&response, rsp_buf).is_err() {
log::warn!(
"server: response PDU exceeds max size for slave {}",
slave_id
);
return None;
}
Some(rsp_buf.split().freeze())
}
#[cfg(test)]
mod tests {
use super::*;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::task::{Context, Poll};
use tokio::io::ReadBuf;
struct NeverQuiet;
impl AsyncRead for NeverQuiet {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let to_fill = buf.remaining().min(16);
let fill_data = [0xAAu8; 16];
buf.put_slice(&fill_data[..to_fill]);
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn drain_stale_data_clean_bus_returns_ok() {
let (mut client, _server) = tokio::io::duplex(64);
let mut scratch = [0u8; 256];
let result = drain_stale_data(&mut client, &mut scratch).await;
assert!(result.is_ok(), "clean bus should return Ok immediately");
}
#[tokio::test]
async fn drain_stale_data_bus_busy_returns_error() {
let mut stream = NeverQuiet;
let mut scratch = [0u8; 256];
let err = drain_stale_data(&mut stream, &mut scratch)
.await
.unwrap_err();
let detail = err.detail();
assert!(
detail.contains("bus busy"),
"expected 'bus busy', got: {detail}"
);
assert!(
matches!(err, ModbusError::Connection(_)),
"bus-busy should be a Connection error (retryable), got: {err:?}"
);
}
#[tokio::test]
async fn drain_stale_data_connection_closed() {
let (mut client, server) = tokio::io::duplex(64);
drop(server); let mut scratch = [0u8; 256];
let err = drain_stale_data(&mut client, &mut scratch)
.await
.unwrap_err();
assert!(
matches!(err, ModbusError::Connection(_)),
"expected Connection error, got: {err:?}"
);
}
#[tokio::test]
async fn rebuild_fails_three_consecutive_times() {
let cfg = ReconnectConfig::new(5, Duration::from_millis(1));
let result = run_with_reconnect(
Some(&cfg),
|| async { Err(ModbusError::connection("test")) },
|| async { false },
ModbusError::connection,
)
.await;
let err = result.unwrap_err();
assert!(
err.detail().contains("3 consecutive times"),
"expected rebuild failure, got: {}",
err.detail()
);
}
#[tokio::test]
async fn rebuild_counter_resets_on_success() {
let rebuild_count = AtomicU32::new(0);
let cfg = ReconnectConfig::new(5, Duration::from_millis(1));
let result = run_with_reconnect(
Some(&cfg),
|| async { Err(ModbusError::connection("test")) },
|| {
let n = rebuild_count.fetch_add(1, Ordering::SeqCst);
async move { n % 2 == 0 } },
ModbusError::connection,
)
.await;
let err = result.unwrap_err();
assert!(
err.detail().contains("max_retries"),
"expected max_retries error (rebuild counter reset), got: {}",
err.detail()
);
}
#[tokio::test]
async fn send_succeeds_after_rebuild() {
let first_call = AtomicBool::new(true);
let cfg = ReconnectConfig::new(3, Duration::from_millis(1));
let result = run_with_reconnect(
Some(&cfg),
|| {
let is_first = first_call.swap(false, Ordering::SeqCst);
async move {
if is_first {
Err(ModbusError::connection("test"))
} else {
Ok(Response::ReadHoldingRegisters(vec![42]))
}
}
},
|| async { true },
ModbusError::connection,
)
.await;
let rsp = result.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
}
#[tokio::test]
async fn max_retries_zero_rebuild_fails_three_times() {
let cfg = ReconnectConfig::new(0, Duration::from_millis(1));
let result = run_with_reconnect(
Some(&cfg),
|| async { Err(ModbusError::connection("test")) },
|| async { false },
ModbusError::connection,
)
.await;
let err = result.unwrap_err();
assert!(
err.detail().contains("3 consecutive times"),
"expected rebuild failure guard, got: {}",
err.detail()
);
}
#[tokio::test]
async fn non_retryable_error_returned_immediately() {
let cfg = ReconnectConfig::new(10, Duration::from_millis(1));
let result = run_with_reconnect(
Some(&cfg),
|| async { Err(ModbusError::protocol("bad CRC")) },
|| async { true },
ModbusError::connection,
)
.await;
let err = result.unwrap_err();
assert!(
matches!(err, ModbusError::Protocol(_)),
"expected Protocol error, got: {err:?}"
);
}
#[tokio::test]
async fn no_reconnect_passthrough() {
let result = run_with_reconnect(
None,
|| async { Ok(Response::ReadHoldingRegisters(vec![99])) },
|| async { true },
ModbusError::connection,
)
.await;
let rsp = result.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![99]));
}
}