1use 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
29const DEFAULT_RTU_T35_US: u64 = 5_000;
32
33type ReconnectFactory<T> = (
35 crate::reconnect::ReconnectConfig,
36 Box<dyn Fn() -> io::Result<T> + Send + Sync>,
37);
38
39pub struct RtuClient<T> {
72 inner: Mutex<RtuInner<T>>,
73 timeout: Duration,
74 reconnect: Option<ReconnectFactory<T>>,
75 bus_timing: Option<Arc<BusTiming>>,
77 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 pub fn new(transport: T) -> Self {
105 Self::with_timeout(transport, Duration::from_secs(5))
106 }
107
108 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 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 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 let mut scratch = [0u8; MAX_ADU_SIZE];
153 send_recv::drain_stale_data(&mut inner.stream, &mut scratch).await?;
154
155 inner.stream.prepare_send().await;
157
158 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 let expected_pdu = expected_rtu_pdu_len(request);
171 let expected_frame = expected_pdu.map(|p| p + 1 + CRC_SIZE); let min_frame: usize = MIN_RTU_FRAME; inner.read_buf.clear();
175 let deadline = Instant::now() + self.timeout;
176 let target = expected_frame.unwrap_or(min_frame);
177
178 send_recv::read_at_least(&mut inner.stream, &mut inner.read_buf, deadline, target).await?;
180
181 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 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]; 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
220fn expected_rtu_pdu_len(request: &Request<'_>) -> Option<usize> {
223 use Request::*;
224 match request {
225 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
243impl<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
297pub 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
317pub struct RtuServer<T> {
338 transport: T,
339 bus_timing: Option<Arc<BusTiming>>,
340 read_timeout: Option<Duration>,
344}
345
346impl<T> RtuServer<T>
347where
348 T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
349{
350 pub fn new(transport: T) -> Self {
352 Self {
353 transport,
354 bus_timing: None,
355 read_timeout: None,
356 }
357 }
358
359 pub fn with_bus_timing(mut self, timing: Arc<BusTiming>) -> Self {
361 self.bus_timing = Some(timing);
362 self
363 }
364
365 pub fn with_read_timeout(mut self, timeout: Duration) -> Self {
372 self.read_timeout = Some(timeout);
373 self
374 }
375
376 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 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 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 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 if try_parse_server_rtu(&buf).is_none() {
419 buf.clear();
420 }
421 continue;
422 }
423 }
424 }
425 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
444fn 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(), ))
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483
484 #[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 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 #[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 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 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 let frame = [0x01, 0x00, 0x00];
629 assert!(try_parse_server_rtu(&frame).is_none());
630 }
631}