use std::fmt::Debug;
use std::io;
use std::sync::Arc;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use bytes::{Buf, Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::Mutex;
use crate::bus_timing::BusTiming;
use crate::client::ModbusClient;
use crate::codec;
use crate::error::ModbusError;
use crate::error::*;
use crate::frame::{Request, Response};
use crate::options::ClientOptions;
use crate::transport::send_recv;
use crate::transport::sniff_io::SniffIo;
use crate::transport::MAX_ADU_SIZE;
use crate::WireTap;
const DEFAULT_ASCII_SERVER_TIMEOUT_US: u64 = 100_000;
type ReconnectFactory<T> = (
crate::reconnect::ReconnectConfig,
Box<dyn Fn() -> io::Result<T> + Send + Sync>,
);
pub struct AsciiClient<T> {
inner: Mutex<AsciiInner<T>>,
timeout: Duration,
reconnect: Option<ReconnectFactory<T>>,
bus_timing: Option<Arc<BusTiming>>,
tap: Option<Arc<dyn WireTap>>,
}
struct AsciiInner<T> {
stream: SniffIo<T>,
write_buf: BytesMut,
read_buf: BytesMut,
}
impl<T> Debug for AsciiClient<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut d = f.debug_struct("AsciiClient");
d.field("timeout", &self.timeout);
if let Some((cfg, _)) = &self.reconnect {
d.field("reconnect_max_retries", &cfg.max_retries());
d.field("reconnect_interval", &cfg.interval());
}
d.finish()
}
}
impl<T> AsciiClient<T>
where
T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
pub fn new(transport: T) -> Self {
Self::with_timeout(transport, Duration::from_secs(5))
}
pub fn with_timeout(transport: T, timeout: Duration) -> Self {
Self {
inner: Mutex::new(AsciiInner {
stream: SniffIo::new(transport, None, None),
write_buf: BytesMut::with_capacity(MAX_ADU_SIZE),
read_buf: BytesMut::with_capacity(MAX_ADU_SIZE),
}),
timeout,
reconnect: None,
bus_timing: None,
tap: None,
}
}
pub fn with_reconnect<F>(mut self, max_retries: u32, backoff: Duration, factory: F) -> Self
where
F: Fn() -> io::Result<T> + Send + Sync + 'static,
{
self.reconnect = Some((
crate::reconnect::ReconnectConfig::new(max_retries, backoff),
Box::new(factory),
));
self
}
async fn send_recv(
&self,
slave_id: u8,
request: &Request<'_>,
) -> Result<Response, ModbusError> {
if *request == Request::Disconnect {
return Err(ModbusError::Other("Disconnected".into()));
}
let mut inner = self.inner.lock().await;
let inner = &mut *inner;
let mut scratch = [0u8; MAX_ADU_SIZE];
send_recv::drain_stale_data(&mut inner.stream, &mut scratch).await?;
inner.stream.prepare_send().await;
send_recv::send_frame(
&mut inner.stream,
&mut inner.write_buf,
slave_id,
self.timeout,
request,
codec::encode_ascii_frame_into,
)
.await?;
inner.read_buf.clear();
let deadline = Instant::now() + self.timeout;
loop {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
if let Some(result) = try_parse_ascii_response(&inner.read_buf, slave_id) {
return result;
}
return Err(ModbusError::timeout("ASCII_RECV_TIMEOUT"));
}
match tokio::time::timeout(remaining, inner.stream.read_buf(&mut inner.read_buf)).await
{
Ok(Ok(0)) => return Err(ModbusError::connection("connection closed")),
Ok(Ok(_)) => {
if let Some(result) = try_parse_ascii_response(&inner.read_buf, slave_id) {
return result;
}
}
Ok(Err(e)) => return Err(ModbusError::from(e)),
Err(_elapsed) => {
if let Some(result) = try_parse_ascii_response(&inner.read_buf, slave_id) {
return result;
}
return Err(ModbusError::timeout("ASCII_RECV_TIMEOUT"));
}
}
}
}
}
const ASCII_START: u8 = b':';
const CR: u8 = b'\r';
const LF: u8 = b'\n';
fn try_parse_ascii_response(buf: &[u8], slave_id: u8) -> Option<Result<Response, ModbusError>> {
let (rsp_slave, pdu, _consumed) = try_parse_ascii(buf)?;
if rsp_slave != slave_id {
return Some(Err(ModbusError::protocol(format!(
"{SLAVE_ID_MISMATCH} {slave_id}, got {rsp_slave}"
))));
}
Some(
Response::try_from(pdu)
.map_err(|e| ModbusError::protocol(format!("{PDU_DECODE_ERROR} {e}"))),
)
}
fn try_parse_ascii(buf: &[u8]) -> Option<(u8, Bytes, usize)> {
let mut offset = 0;
loop {
let remaining = buf.get(offset..)?;
let crlf_pos = remaining.windows(2).position(|w| w == [CR, LF])?;
let line_end = offset + crlf_pos + 2; let content = &buf[offset..line_end - 2];
if !content.is_empty() && content[0] == ASCII_START {
let hex_chars = &content[1..];
if hex_chars.len() % 2 == 0 {
let mut bytes = Vec::with_capacity(hex_chars.len() / 2);
let mut ok = true;
for chunk in hex_chars.chunks(2) {
match hex_decode_byte(chunk[0], chunk[1]) {
Ok(b) => bytes.push(b),
Err(_) => {
ok = false;
break;
}
}
}
if ok && bytes.len() >= 2 {
let (data, lrc_byte) = bytes.split_at(bytes.len() - 1);
if lrc_byte[0] == codec::calculate_lrc(data) && data.len() >= 2 {
return Some((data[0], Bytes::copy_from_slice(&data[1..]), line_end));
}
}
}
}
offset = line_end;
}
}
fn hex_decode_byte(hi: u8, lo: u8) -> io::Result<u8> {
fn nibble(c: u8) -> io::Result<u8> {
match c {
b'0'..=b'9' => Ok(c - b'0'),
b'A'..=b'F' => Ok(c - b'A' + 10),
b'a'..=b'f' => Ok(c - b'a' + 10),
_ => Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid hex character",
)),
}
}
Ok((nibble(hi)? << 4) | nibble(lo)?)
}
impl<T> AsciiClient<SniffIo<T>>
where
T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
pub fn with_reconnect_on<F>(mut self, max_retries: u32, backoff: Duration, factory: F) -> Self
where
F: Fn() -> io::Result<T> + Send + Sync + 'static,
{
let timing = self.bus_timing.clone();
let tap = self.tap.clone();
self.reconnect = Some((
crate::reconnect::ReconnectConfig::new(max_retries, backoff),
Box::new(move || {
let raw = factory()?;
Ok(SniffIo::new(raw, tap.clone(), timing.clone()))
}),
));
self
}
}
pub fn with_options<T: AsyncRead + AsyncWrite + Unpin + Send + 'static>(
transport: T,
opts: ClientOptions,
) -> AsciiClient<SniffIo<T>> {
let tap = opts.tap().cloned();
let timing = opts.bus_timing.clone();
let mut stream = SniffIo::new(transport, tap, timing.clone());
if let Some(cap) = opts.data_channel_capacity {
stream = stream.with_channel_capacity(cap);
}
if let Some(cap) = opts.tap_channel_capacity {
stream = stream.with_tap_channel_capacity(cap);
}
let mut raw = AsciiClient::with_timeout(stream, opts.timeout);
raw.bus_timing = timing;
raw.tap = opts.tap().cloned();
raw
}
#[async_trait]
impl<T> ModbusClient for AsciiClient<T>
where
T: AsyncRead + AsyncWrite + Send + Unpin + 'static,
{
async fn call(&self, slave: u8, request: Request<'_>) -> Result<Response, ModbusError> {
let request = request.into_owned();
let slave_id = slave;
let timing = self.bus_timing.clone();
send_recv::run_with_reconnect(
self.reconnect.as_ref().map(|(cfg, _)| cfg),
|| self.send_recv(slave_id, &request),
|| async {
let mut inner = self.inner.lock().await;
if let Some((_, factory)) = &self.reconnect {
if let Ok(new) = factory() {
inner.stream = SniffIo::new(new, self.tap.clone(), timing.clone());
inner.write_buf.clear();
inner.read_buf.clear();
return true;
}
}
false
},
ModbusError::serial,
)
.await
}
}
pub struct AsciiServer<T> {
transport: T,
bus_timing: Option<Arc<BusTiming>>,
read_timeout: Option<Duration>,
}
impl<T> AsciiServer<T>
where
T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
pub fn new(transport: T) -> Self {
Self {
transport,
bus_timing: None,
read_timeout: None,
}
}
pub fn with_bus_timing(mut self, timing: Arc<BusTiming>) -> Self {
self.bus_timing = Some(timing);
self
}
pub fn with_read_timeout(mut self, timeout: Duration) -> Self {
self.read_timeout = Some(timeout);
self
}
pub async fn serve_forever<S>(self, service: S) -> std::io::Result<()>
where
S: crate::server::Service + Send + Sync + 'static,
{
let timing = self.bus_timing.clone();
let mut stream = SniffIo::new(self.transport, None, timing.clone());
let mut buf = BytesMut::with_capacity(MAX_ADU_SIZE);
let mut rsp_buf = BytesMut::with_capacity(MAX_ADU_SIZE);
let mut frame_buf = BytesMut::with_capacity(MAX_ADU_SIZE);
let read_timeout = self
.read_timeout
.or_else(|| timing.as_ref().map(|t| t.min_spacing()))
.unwrap_or(Duration::from_micros(DEFAULT_ASCII_SERVER_TIMEOUT_US));
loop {
let mut tmp = [0u8; MAX_ADU_SIZE];
if buf.is_empty() {
match stream.read(&mut tmp).await {
Ok(0) => break,
Ok(n) => buf.extend_from_slice(&tmp[..n]),
Err(_) => break,
}
} else {
match tokio::time::timeout(read_timeout, stream.read(&mut tmp)).await {
Ok(Ok(0)) => break,
Ok(Ok(n)) => buf.extend_from_slice(&tmp[..n]),
Ok(Err(_)) => break,
Err(_elapsed) => {
buf.clear();
continue;
}
}
}
while let Some((slave_id, pdu, consumed)) = try_parse_ascii(&buf) {
buf.advance(consumed);
if let Some(rsp_data) =
send_recv::process_server_request(&pdu, slave_id, &service, &mut rsp_buf).await
{
stream.prepare_send().await;
frame_buf.clear();
codec::encode_ascii_frame_into(&rsp_data, &mut frame_buf);
if stream.write_all(&frame_buf).await.is_err() {
return Ok(());
}
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_ascii_frame(bytes: &[u8]) -> Vec<u8> {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
let lrc = bytes
.iter()
.fold(0u8, |acc, &b| acc.wrapping_add(b))
.wrapping_neg();
let mut frame = vec![b':'];
for &b in bytes.iter().chain(std::iter::once(&lrc)) {
frame.push(HEX[(b >> 4) as usize]);
frame.push(HEX[(b & 0x0F) as usize]);
}
frame.extend_from_slice(b"\r\n");
frame
}
#[test]
fn empty_buffer() {
assert!(try_parse_ascii(&[]).is_none());
}
#[test]
fn no_crlf_terminator() {
let data = b":010300000001FB\r";
assert!(try_parse_ascii(data).is_none());
}
#[test]
fn cr_without_lf() {
let data = b":010300000001FB\rX";
assert!(try_parse_ascii(data).is_none());
}
#[test]
fn valid_read_holding_request() {
let frame = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let (slave, pdu, consumed) = try_parse_ascii(&frame).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
assert_eq!(pdu.len(), 5);
assert_eq!(consumed, frame.len());
}
#[test]
fn no_colon_prefix_invalid() {
let data = b"010300000001FB\r\n";
assert!(try_parse_ascii(data).is_none());
}
#[test]
fn bad_hex_characters() {
let data = b":01G300000001XX\r\n";
assert!(try_parse_ascii(data).is_none());
}
#[test]
fn odd_hex_length() {
let valid = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let mut data = b":01030\r\n".to_vec();
data.extend_from_slice(&valid);
let (slave, pdu, _) = try_parse_ascii(&data).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
}
#[test]
fn bad_lrc_recovery() {
let mut bad = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let crlf_pos = bad.windows(2).position(|w| w == b"\r\n").unwrap();
bad[crlf_pos - 2] = b'F';
bad[crlf_pos - 1] = b'F';
let good = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let mut data = bad;
data.extend_from_slice(&good);
let (slave, pdu, consumed) = try_parse_ascii(&data).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
assert!(consumed > 20);
}
#[test]
fn too_short_after_hex_decode() {
let valid = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let mut data = b":01\r\n".to_vec();
data.extend_from_slice(&valid);
let (slave, pdu, _) = try_parse_ascii(&data).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
}
#[test]
fn lrc_validates_correctly() {
let frame = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let (slave, pdu, _) = try_parse_ascii(&frame).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
}
#[test]
fn lowercase_hex_accepted() {
let frame = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let lower: Vec<u8> = frame
.iter()
.map(|&b| if b.is_ascii_uppercase() { b + 32 } else { b })
.collect();
let (slave, pdu, _) = try_parse_ascii(&lower).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[0], 3);
}
#[test]
fn multiple_valid_frames_parses_first_only() {
let first = make_ascii_frame(&[0x01, 0x03, 0x00, 0x00, 0x00, 0x01]);
let second = make_ascii_frame(&[0x01, 0x03, 0x00, 0x01, 0x00, 0x01]);
let mut data = first.clone();
data.extend_from_slice(&second);
let (slave, pdu, consumed) = try_parse_ascii(&data).unwrap();
assert_eq!(slave, 1);
assert_eq!(pdu[3], 0);
assert_eq!(consumed, first.len());
}
#[test]
fn hex_decode_all_nibbles() {
assert_eq!(hex_decode_byte(b'0', b'0').unwrap(), 0x00);
assert_eq!(hex_decode_byte(b'0', b'1').unwrap(), 0x01);
assert_eq!(hex_decode_byte(b'1', b'0').unwrap(), 0x10);
assert_eq!(hex_decode_byte(b'F', b'F').unwrap(), 0xFF);
assert_eq!(hex_decode_byte(b'f', b'f').unwrap(), 0xFF);
assert_eq!(hex_decode_byte(b'A', b'a').unwrap(), 0xAA);
}
#[test]
fn hex_decode_invalid_high_nibble() {
assert!(hex_decode_byte(b'G', b'0').is_err());
}
#[test]
fn hex_decode_invalid_low_nibble() {
assert!(hex_decode_byte(b'0', b'G').is_err());
}
}