Skip to main content

oms_modbus/transport/
rtu.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2//! Modbus RTU client and server — generic over transport.
3//!
4//! Works with any `AsyncRead + AsyncWrite` byte transport, including
5//! serial ports and TCP (RTU over TCP).
6
7use std::fmt::Debug;
8use std::io;
9use std::sync::Arc;
10use std::time::{Duration, Instant};
11
12use async_trait::async_trait;
13use bytes::{Buf, Bytes, BytesMut};
14use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
15use tokio::sync::Mutex;
16
17use crate::bus_timing::BusTiming;
18use crate::client::ModbusClient;
19use crate::codec;
20use crate::error::ModbusError;
21use crate::error::*;
22use crate::frame::{Request, Response};
23use crate::options::ClientOptions;
24use crate::transport::send_recv;
25use crate::transport::sniff_io::SniffIo;
26use crate::transport::{CRC_SIZE, MAX_ADU_SIZE, MIN_RTU_FRAME};
27use crate::WireTap;
28
29/// Default T3.5 inter-frame timeout for RTU server when `BusTiming` is not
30/// configured. 5 ms ≈ T3.5 @ 9600 baud (4.01 ms), rounded up for margin.
31const DEFAULT_RTU_T35_US: u64 = 5_000;
32
33/// Reconnect factory for direct transport.
34type ReconnectFactory<T> = (
35    crate::reconnect::ReconnectConfig,
36    Box<dyn Fn() -> io::Result<T> + Send + Sync>,
37);
38
39/// Modbus RTU client — `Send + Sync + Clone` via internal `Mutex`.
40///
41/// Call `.with_reconnect(max, backoff, factory)` to enable auto-reconnect
42/// on transport errors (e.g. USB serial unplugged). The factory is called to
43/// create a new transport instance on each reconnect attempt.
44///
45/// For bus-level capture, use [`with_options`] which wraps the transport in
46/// [`SniffIo`] automatically.
47///
48/// # Example
49///
50/// ```no_run
51/// use std::sync::Arc;
52/// use std::time::Duration;
53/// use oms_modbus::*;
54///
55/// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
56/// // In-memory duplex — works with serial, TCP, Unix sockets too
57/// let (client_port, server_port) = tokio::io::duplex(1024);
58///
59/// // With capture + timing via ClientOptions
60/// let cap = Arc::new(BusCapture::unbounded());
61/// let opts = ClientOptions::default()
62///     .with_timeout(Duration::from_secs(3))
63///     .with_tap(cap.clone())
64///     .with_bus_timing(BusTiming::rtu_35t(9600));
65/// let client = rtu::with_options(client_port, opts);
66///
67/// let regs = client.read_holding_registers(1, 0, 5).await?;
68/// # Ok(())
69/// # }
70/// ```
71pub struct RtuClient<T> {
72    inner: Mutex<RtuInner<T>>,
73    timeout: Duration,
74    reconnect: Option<ReconnectFactory<T>>,
75    /// Preserved for reconnect rebuilds — active timing lives in SniffIo.
76    bus_timing: Option<Arc<BusTiming>>,
77    /// Preserved for reconnect rebuilds — active tap lives in SniffIo.
78    tap: Option<Arc<dyn WireTap>>,
79}
80
81struct RtuInner<T> {
82    stream: SniffIo<T>,
83    write_buf: BytesMut,
84    read_buf: BytesMut,
85}
86
87impl<T> Debug for RtuClient<T> {
88    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
89        let mut d = f.debug_struct("RtuClient");
90        d.field("timeout", &self.timeout);
91        if let Some((cfg, _)) = &self.reconnect {
92            d.field("reconnect_max_retries", &cfg.max_retries());
93            d.field("reconnect_interval", &cfg.interval());
94        }
95        d.finish()
96    }
97}
98
99impl<T> RtuClient<T>
100where
101    T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
102{
103    /// Create an RTU client with a 5-second default timeout.
104    pub fn new(transport: T) -> Self {
105        Self::with_timeout(transport, Duration::from_secs(5))
106    }
107
108    /// Create an RTU client with an explicit request timeout.
109    pub fn with_timeout(transport: T, timeout: Duration) -> Self {
110        Self {
111            inner: Mutex::new(RtuInner {
112                stream: SniffIo::new(transport, None, None),
113                write_buf: BytesMut::with_capacity(MAX_ADU_SIZE),
114                read_buf: BytesMut::with_capacity(MAX_ADU_SIZE),
115            }),
116            timeout,
117            reconnect: None,
118            bus_timing: None,
119            tap: None,
120        }
121    }
122
123    /// Enable auto-reconnect with a factory that creates the FULL transport.
124    ///
125    /// For `RtuClient<SniffIo<T>>`, prefer `with_reconnect_on` which
126    /// accepts a raw-transport factory and handles SniffIo wrapping.
127    pub fn with_reconnect<F>(mut self, max_retries: u32, backoff: Duration, factory: F) -> Self
128    where
129        F: Fn() -> io::Result<T> + Send + Sync + 'static,
130    {
131        self.reconnect = Some((
132            crate::reconnect::ReconnectConfig::new(max_retries, backoff),
133            Box::new(factory),
134        ));
135        self
136    }
137
138    /// ── RTU receive: length-based ────────────────────────────────────────
139    ///
140    /// The client knows what request it sent → it knows the expected response
141    /// frame length. Reads exactly that many bytes (with timeout), then
142    /// validates CRC once. No per-byte CRC scanning.
143    async fn send_recv(
144        &self,
145        slave_id: u8,
146        request: &Request<'_>,
147    ) -> Result<Response, ModbusError> {
148        let mut inner = self.inner.lock().await;
149        let inner = &mut *inner;
150
151        // ── Pre-send: drain stale data (all protocols, half-duplex) ──────
152        let mut scratch = [0u8; MAX_ADU_SIZE];
153        send_recv::drain_stale_data(&mut inner.stream, &mut scratch).await?;
154
155        // ── Prepare bus timing ────────────────────────────────────────────
156        inner.stream.prepare_send().await;
157
158        // ── Send ──────────────────────────────────────────────────────────
159        send_recv::send_frame(
160            &mut inner.stream,
161            &mut inner.write_buf,
162            slave_id,
163            self.timeout,
164            request,
165            codec::encode_rtu_frame,
166        )
167        .await?;
168
169        // ── Receive: read expected frame length ───────────────────────────
170        let expected_pdu = expected_rtu_pdu_len(request);
171        let expected_frame = expected_pdu.map(|p| p + 1 + CRC_SIZE); // slave_id(1) + pdu + crc
172        let min_frame: usize = MIN_RTU_FRAME; // slave_id(1) + fc|0x80(1) + code(1) + crc(2)
173
174        inner.read_buf.clear();
175        let deadline = Instant::now() + self.timeout;
176        let target = expected_frame.unwrap_or(min_frame);
177
178        // Read until we have `target` bytes or timeout
179        send_recv::read_at_least(&mut inner.stream, &mut inner.read_buf, deadline, target).await?;
180
181        // Timeout may leave us with fewer than target bytes. Validate CRC
182        // on whatever we have (≥5 bytes minimum).
183        if inner.read_buf.len() < min_frame {
184            let got = inner.read_buf.len();
185            if got == 0 {
186                return Err(ModbusError::timeout(RTU_RECV_TIMEOUT));
187            }
188            return Err(ModbusError::timeout(format!(
189                "RTU short frame ({got} bytes, need at least {min_frame})"
190            )));
191        }
192
193        // Validate CRC on accumulated data
194        let buf = &inner.read_buf;
195        let crc_received = u16::from_le_bytes([buf[buf.len() - 2], buf[buf.len() - 1]]);
196        let crc_calc = codec::calculate_crc(&buf[..buf.len() - 2]);
197        if crc_received != crc_calc {
198            return Err(ModbusError::protocol("RTU CRC mismatch"));
199        }
200
201        let frame = &buf[..buf.len() - 2]; // strip CRC
202        if frame.len() < 2 {
203            return Err(ModbusError::protocol("RTU frame too short"));
204        }
205
206        let rsp_slave = frame[0];
207        let pdu = Bytes::copy_from_slice(&frame[1..]);
208
209        if rsp_slave != slave_id {
210            return Err(ModbusError::protocol(format!(
211                "{SLAVE_ID_MISMATCH} {slave_id}, got {rsp_slave}"
212            )));
213        }
214
215        Response::try_from(pdu)
216            .map_err(|e| ModbusError::protocol(format!("{PDU_DECODE_ERROR} {e}")))
217    }
218}
219
220// ── Expected RTU response PDU length ────────────────────────────────────
221
222fn expected_rtu_pdu_len(request: &Request<'_>) -> Option<usize> {
223    use Request::*;
224    match request {
225        // PDU for read responses: FC(1) + byte_count(1) + data
226        ReadCoils(_, count) | ReadDiscreteInputs(_, count) => {
227            Some(2 + (*count as usize).div_ceil(8))
228        }
229        ReadHoldingRegisters(_, count) | ReadInputRegisters(_, count) => {
230            Some(2 + *count as usize * 2)
231        }
232        WriteSingleCoil(_, _)
233        | WriteSingleRegister(_, _)
234        | WriteMultipleCoils(_, _)
235        | WriteMultipleRegisters(_, _) => Some(5),
236        ReadWriteMultipleRegisters(_, count, _, _) => Some(2 + *count as usize * 2),
237        MaskWriteRegister(_, _, _) => Some(7),
238        Diagnostic(_, _) => Some(5),
239        Disconnect => None,
240    }
241}
242
243// ── Reconnect for SniffIo-wrapped clients ───────────────────────────────
244
245impl<T> RtuClient<SniffIo<T>>
246where
247    T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
248{
249    pub fn with_reconnect_on<F>(mut self, max_retries: u32, backoff: Duration, factory: F) -> Self
250    where
251        F: Fn() -> io::Result<T> + Send + Sync + 'static,
252    {
253        let timing = self.bus_timing.clone();
254        let tap = self.tap.clone();
255        self.reconnect = Some((
256            crate::reconnect::ReconnectConfig::new(max_retries, backoff),
257            Box::new(move || {
258                let raw = factory()?;
259                Ok(SniffIo::new(raw, tap.clone(), timing.clone()))
260            }),
261        ));
262        self
263    }
264}
265
266#[async_trait]
267impl<T> ModbusClient for RtuClient<T>
268where
269    T: AsyncRead + AsyncWrite + Send + Unpin + 'static,
270{
271    async fn call(&self, slave: u8, request: Request<'_>) -> Result<Response, ModbusError> {
272        let request = request.into_owned();
273        let slave_id = slave;
274        let timing = self.bus_timing.clone();
275
276        send_recv::run_with_reconnect(
277            self.reconnect.as_ref().map(|(cfg, _)| cfg),
278            || self.send_recv(slave_id, &request),
279            || async {
280                let mut inner = self.inner.lock().await;
281                if let Some((_, factory)) = &self.reconnect {
282                    if let Ok(new) = factory() {
283                        inner.stream = SniffIo::new(new, self.tap.clone(), timing.clone());
284                        inner.write_buf.clear();
285                        inner.read_buf.clear();
286                        return true;
287                    }
288                }
289                false
290            },
291            ModbusError::serial,
292        )
293        .await
294    }
295}
296
297/// Create an RTU client with optional capture and reconnect via SniffIo.
298pub fn with_options<T: AsyncRead + AsyncWrite + Unpin + Send + 'static>(
299    transport: T,
300    opts: ClientOptions,
301) -> RtuClient<SniffIo<T>> {
302    let tap = opts.tap().cloned();
303    let timing = opts.bus_timing.clone();
304    let mut stream = SniffIo::new(transport, tap, timing.clone());
305    if let Some(cap) = opts.data_channel_capacity {
306        stream = stream.with_channel_capacity(cap);
307    }
308    if let Some(cap) = opts.tap_channel_capacity {
309        stream = stream.with_tap_channel_capacity(cap);
310    }
311    let mut raw = RtuClient::with_timeout(stream, opts.timeout);
312    raw.bus_timing = timing;
313    raw.tap = opts.tap().cloned();
314    raw
315}
316
317// ── RTU Server ──────────────────────────────────────────────────────────
318
319/// Modbus RTU server over any byte transport (serial, DuplexStream, etc.).
320///
321/// # Example
322///
323/// ```no_run
324/// use std::sync::Arc;
325/// use oms_modbus::*;
326///
327/// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
328/// let (port, _server_side) = tokio::io::duplex(64);
329/// let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
330///
331/// let server = rtu::RtuServer::new(port)
332///     .with_bus_timing(Arc::new(BusTiming::rtu_35t(9600)));
333/// tokio::spawn(async move { server.serve_forever(store).await.ok(); });
334/// # Ok(())
335/// # }
336/// ```
337pub struct RtuServer<T> {
338    transport: T,
339    bus_timing: Option<Arc<BusTiming>>,
340    /// Per-frame read timeout — when the bus is silent for this long,
341    /// accumulated bytes are flushed as a broken frame. Defaults to
342    /// `BusTiming::min_spacing()` (T3.5) when set, or 5 ms otherwise.
343    read_timeout: Option<Duration>,
344}
345
346impl<T> RtuServer<T>
347where
348    T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
349{
350    /// Create an RTU server on the given byte transport.
351    pub fn new(transport: T) -> Self {
352        Self {
353            transport,
354            bus_timing: None,
355            read_timeout: None,
356        }
357    }
358
359    /// Enforce bus timing on the server side (optional).
360    pub fn with_bus_timing(mut self, timing: Arc<BusTiming>) -> Self {
361        self.bus_timing = Some(timing);
362        self
363    }
364
365    /// Set a custom per-frame read timeout for detecting broken/partial
366    /// frames. When this timeout expires without receiving new bytes,
367    /// any unparseable data in the buffer is discarded.
368    ///
369    /// Default: [`BusTiming::min_spacing`] if bus timing is configured,
370    /// otherwise 5 ms (≈ T3.5 @ 9600 baud).
371    pub fn with_read_timeout(mut self, timeout: Duration) -> Self {
372        self.read_timeout = Some(timeout);
373        self
374    }
375
376    /// Process requests indefinitely until the transport closes.
377    ///
378    /// Frame boundaries are detected by inter-byte timeout (3.5T silence).
379    /// Bytes are accumulated until the bus goes quiet, then CRC is validated.
380    pub async fn serve_forever<S>(self, service: S) -> std::io::Result<()>
381    where
382        S: crate::server::Service + Send + Sync + 'static,
383    {
384        let timing = self.bus_timing.clone();
385        let mut stream = SniffIo::new(self.transport, None, timing.clone());
386        let mut buf = BytesMut::with_capacity(MAX_ADU_SIZE);
387        let mut rsp_buf = BytesMut::with_capacity(MAX_ADU_SIZE);
388        let mut frame_buf = BytesMut::with_capacity(MAX_ADU_SIZE);
389
390        // Resolve read timeout: explicit user setting > BusTiming > default.
391        let read_timeout = self
392            .read_timeout
393            .or_else(|| timing.as_ref().map(|t| t.min_spacing()))
394            .unwrap_or(Duration::from_micros(DEFAULT_RTU_T35_US));
395
396        loop {
397            let mut tmp = [0u8; MAX_ADU_SIZE];
398            if buf.is_empty() {
399                // No pending data — wait indefinitely for the next frame.
400                match stream.read(&mut tmp).await {
401                    Ok(0) => break,
402                    Ok(n) => buf.extend_from_slice(&tmp[..n]),
403                    Err(_) => break,
404                }
405            } else {
406                // Have accumulated bytes — wait T3.5 for more, then decide
407                // whether the frame is complete or broken.
408                match tokio::time::timeout(read_timeout, stream.read(&mut tmp)).await {
409                    Ok(Ok(0)) => break,
410                    Ok(Ok(n)) => buf.extend_from_slice(&tmp[..n]),
411                    Ok(Err(_)) => break,
412                    Err(_elapsed) => {
413                        // T3.5 silence — the accumulated bytes are a
414                        // complete (or broken) frame. Try one last parse
415                        // before discarding; if CRC matches, process it.
416                        // Otherwise clear and wait for the next clean frame
417                        // (libmodbus flush equivalent).
418                        if try_parse_server_rtu(&buf).is_none() {
419                            buf.clear();
420                        }
421                        continue;
422                    }
423                }
424            }
425            // Try to parse complete frames from buf.
426            while let Some((slave_id, pdu, consumed)) = try_parse_server_rtu(&buf) {
427                buf.advance(consumed);
428                if let Some(rsp_data) =
429                    send_recv::process_server_request(&pdu, slave_id, &service, &mut rsp_buf).await
430                {
431                    stream.prepare_send().await;
432                    frame_buf.clear();
433                    codec::encode_rtu_frame(&rsp_data, &mut frame_buf);
434                    if stream.write_all(&frame_buf).await.is_err() {
435                        return Ok(());
436                    }
437                }
438            }
439        }
440        Ok(())
441    }
442}
443
444/// Server-side RTU frame parsing: validate CRC, check FC validity.
445///
446/// Returns `(slave_id, pdu_bytes, total_consumed)` for a valid frame, or
447/// `None` if the buffer doesn't contain a recognizable Modbus RTU frame.
448///
449/// Unknown function codes are silently discarded (no Exception 0x01
450/// response). This is a deliberate choice: on noisy RS-485 buses, a CRC-16
451/// collision with random bytes that happen to decode as an unknown FC would
452/// generate spurious exception responses, adding unnecessary bus traffic.
453/// Many industrial Modbus devices follow the same practice.  Configure
454/// `read_timeout` to control how long the server waits before flushing
455/// unparseable bytes via the T3.5 silence window.
456fn try_parse_server_rtu(buf: &[u8]) -> Option<(u8, Bytes, usize)> {
457    if buf.len() < MIN_RTU_FRAME {
458        return None;
459    }
460    let crc_received = u16::from_le_bytes([buf[buf.len() - 2], buf[buf.len() - 1]]);
461    let crc_calc = codec::calculate_crc(&buf[..buf.len() - 2]);
462    if crc_received != crc_calc {
463        return None;
464    }
465    let frame = &buf[..buf.len() - 2];
466    if frame.len() < 2 {
467        return None;
468    }
469    let fc = frame[1];
470    if !crate::frame::is_known_function_code(fc) {
471        return None;
472    }
473    Some((
474        frame[0],
475        Bytes::copy_from_slice(&frame[1..]),
476        buf.len(), // consumed = frame + crc
477    ))
478}
479
480#[cfg(test)]
481mod tests {
482    use super::*;
483
484    // ── expected_rtu_pdu_len ────────────────────────────────────────────
485
486    #[test]
487    fn read_coils_pdu_len() {
488        assert_eq!(expected_rtu_pdu_len(&Request::ReadCoils(1, 1)), Some(3));
489        assert_eq!(expected_rtu_pdu_len(&Request::ReadCoils(1, 8)), Some(3));
490        assert_eq!(expected_rtu_pdu_len(&Request::ReadCoils(1, 9)), Some(4));
491    }
492
493    #[test]
494    fn read_discrete_inputs_pdu_len() {
495        assert_eq!(
496            expected_rtu_pdu_len(&Request::ReadDiscreteInputs(1, 2000)),
497            Some(2 + 250)
498        );
499    }
500
501    #[test]
502    fn read_holding_registers_pdu_len() {
503        assert_eq!(
504            expected_rtu_pdu_len(&Request::ReadHoldingRegisters(1, 1)),
505            Some(4)
506        );
507        assert_eq!(
508            expected_rtu_pdu_len(&Request::ReadHoldingRegisters(1, 125)),
509            Some(252)
510        );
511    }
512
513    #[test]
514    fn read_input_registers_pdu_len() {
515        assert_eq!(
516            expected_rtu_pdu_len(&Request::ReadInputRegisters(1, 125)),
517            Some(252)
518        );
519    }
520
521    #[test]
522    fn write_single_coil_pdu_len() {
523        assert_eq!(
524            expected_rtu_pdu_len(&Request::WriteSingleCoil(1, true)),
525            Some(5)
526        );
527    }
528
529    #[test]
530    fn write_single_register_pdu_len() {
531        assert_eq!(
532            expected_rtu_pdu_len(&Request::WriteSingleRegister(1, 0x1234)),
533            Some(5)
534        );
535    }
536
537    #[test]
538    fn write_multiple_coils_pdu_len() {
539        let coils: std::borrow::Cow<'_, [bool]> = std::borrow::Cow::Owned(vec![true; 10]);
540        assert_eq!(
541            expected_rtu_pdu_len(&Request::WriteMultipleCoils(1, coils)),
542            Some(5)
543        );
544    }
545
546    #[test]
547    fn write_multiple_registers_pdu_len() {
548        let regs: std::borrow::Cow<'_, [u16]> = std::borrow::Cow::Owned(vec![0; 10]);
549        assert_eq!(
550            expected_rtu_pdu_len(&Request::WriteMultipleRegisters(1, regs)),
551            Some(5)
552        );
553    }
554
555    #[test]
556    fn read_write_multiple_registers_pdu_len() {
557        let regs: std::borrow::Cow<'_, [u16]> = std::borrow::Cow::Owned(vec![0; 3]);
558        assert_eq!(
559            expected_rtu_pdu_len(&Request::ReadWriteMultipleRegisters(1, 3, 1, regs)),
560            Some(8)
561        );
562    }
563
564    #[test]
565    fn mask_write_register_pdu_len() {
566        assert_eq!(
567            expected_rtu_pdu_len(&Request::MaskWriteRegister(1, 0x0000, 0xFFFF)),
568            Some(7)
569        );
570    }
571
572    #[test]
573    fn diagnostic_pdu_len() {
574        // Diagnostic(sub_function, data) — both u16
575        assert_eq!(
576            expected_rtu_pdu_len(&Request::Diagnostic(0x0000, 0x0000)),
577            Some(5)
578        );
579    }
580
581    #[test]
582    fn disconnect_pdu_len_is_none() {
583        assert_eq!(expected_rtu_pdu_len(&Request::Disconnect), None);
584    }
585
586    // ── try_parse_server_rtu ────────────────────────────────────────────
587
588    #[test]
589    fn server_parse_empty_buffer() {
590        assert!(try_parse_server_rtu(&[]).is_none());
591    }
592
593    #[test]
594    fn server_parse_too_short() {
595        assert!(try_parse_server_rtu(&[0x01, 0x03, 0x00, 0x00]).is_none());
596    }
597
598    #[test]
599    fn server_parse_valid_read_holding_response() {
600        // slave=1, FC=3, byte_count=2, data=[0, 0]
601        let mut frame = vec![0x01, 0x03, 0x02, 0x00, 0x00];
602        let crc = crate::codec::calculate_crc(&frame);
603        frame.extend_from_slice(&crc.to_le_bytes());
604        let (slave, pdu, consumed) = try_parse_server_rtu(&frame).unwrap();
605        assert_eq!(slave, 1);
606        assert_eq!(pdu[0], 0x03);
607        assert_eq!(consumed, frame.len());
608    }
609
610    #[test]
611    fn server_parse_crc_mismatch() {
612        let frame = [0x01, 0x03, 0x00, 0x00, 0x00, 0x01, 0xFF, 0xFF];
613        assert!(try_parse_server_rtu(&frame).is_none());
614    }
615
616    #[test]
617    fn server_parse_unknown_function_code_rejected() {
618        // Valid CRC but FC=0x46 (unknown) → silently discarded
619        let mut frame = vec![0x01, 0x46, 0x00, 0x00, 0x01];
620        let crc = crate::codec::calculate_crc(&frame);
621        frame.extend_from_slice(&crc.to_le_bytes());
622        assert!(try_parse_server_rtu(&frame).is_none());
623    }
624
625    #[test]
626    fn server_parse_zero_length_frame_after_crc_removal() {
627        // Frame with only slave_id + CRC → after CRC strip, < 2 bytes
628        let frame = [0x01, 0x00, 0x00];
629        assert!(try_parse_server_rtu(&frame).is_none());
630    }
631}