use std::sync::atomic::{AtomicU32, AtomicU8, Ordering};
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use oms_modbus::*;
struct NoopHook;
#[async_trait]
impl ServerHook for NoopHook {}
struct ShortCircuitHook {
fixed_value: u16,
call_count: AtomicU8,
}
#[async_trait]
impl ServerHook for ShortCircuitHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::ReadHoldingRegisters(_addr, qty) = request {
self.call_count.fetch_add(1, Ordering::Relaxed);
let regs = vec![self.fixed_value; *qty as usize];
Some(Response::ReadHoldingRegisters(regs))
} else {
None
}
}
}
struct SuppressExceptionHook;
#[async_trait]
impl ServerHook for SuppressExceptionHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Ok(rsp) => Ok(rsp),
Err(_) => Ok(Response::ReadHoldingRegisters(vec![0; 1])),
}
}
}
struct LoggingHook {
last_slave: AtomicU8,
before_count: AtomicU8,
after_count: AtomicU8,
}
impl LoggingHook {
fn new() -> Self {
Self {
last_slave: AtomicU8::new(0),
before_count: AtomicU8::new(0),
after_count: AtomicU8::new(0),
}
}
}
#[async_trait]
impl ServerHook for LoggingHook {
async fn before_call(&self, slave: u8, _request: &Request<'_>) -> Option<Response> {
self.last_slave.store(slave, Ordering::Relaxed);
self.before_count.fetch_add(1, Ordering::Relaxed);
None
}
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
self.after_count.fetch_add(1, Ordering::Relaxed);
result
}
}
struct AllReadsHook {
coil_val: bool,
reg_val: u16,
}
#[async_trait]
impl ServerHook for AllReadsHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::ReadCoils(_addr, qty) => {
Some(Response::ReadCoils(vec![self.coil_val; *qty as usize]))
}
Request::ReadDiscreteInputs(_addr, qty) => Some(Response::ReadDiscreteInputs(vec![
self.coil_val;
*qty as usize
])),
Request::ReadHoldingRegisters(_addr, qty) => {
Some(Response::ReadHoldingRegisters(vec![
self.reg_val;
*qty as usize
]))
}
Request::ReadInputRegisters(_addr, qty) => Some(Response::ReadInputRegisters(vec![
self.reg_val;
*qty as usize
])),
_ => None,
}
}
}
#[tokio::test]
async fn before_call_short_circuits_read_coils() {
let store = Arc::new(SlaveStore::new());
let hook = AllReadsHook {
coil_val: true,
reg_val: 0,
};
let hooked = HookedService::new(store, hook);
let rsp = hooked.call(Request::ReadCoils(0, 8)).await.unwrap();
assert_eq!(rsp, Response::ReadCoils(vec![true; 8]));
}
#[tokio::test]
async fn before_call_short_circuits_read_discrete_inputs() {
let store = Arc::new(SlaveStore::new());
let hook = AllReadsHook {
coil_val: false,
reg_val: 0,
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::ReadDiscreteInputs(10, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadDiscreteInputs(vec![false; 3]));
}
#[tokio::test]
async fn before_call_short_circuits_read_input_registers() {
let store = Arc::new(SlaveStore::new());
let hook = AllReadsHook {
coil_val: false,
reg_val: 777,
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::ReadInputRegisters(0, 4))
.await
.unwrap();
assert_eq!(rsp, Response::ReadInputRegisters(vec![777; 4]));
}
struct WriteEchoHook {
last_written: Mutex<Vec<(u16, u16)>>,
}
impl WriteEchoHook {
fn new() -> Self {
Self {
last_written: Mutex::new(Vec::new()),
}
}
}
#[async_trait]
impl ServerHook for WriteEchoHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::WriteSingleRegister(addr, val) => {
self.last_written.lock().unwrap().push((*addr, *val));
Some(Response::WriteSingleRegister(*addr, *val))
}
Request::WriteMultipleRegisters(addr, vals) => {
let vals: Vec<u16> = vals.to_vec();
self.last_written
.lock()
.unwrap()
.extend(vals.iter().enumerate().map(|(i, &v)| (*addr + i as u16, v)));
Some(Response::WriteMultipleRegisters(*addr, vals.len() as u16))
}
_ => None,
}
}
}
#[tokio::test]
async fn before_call_short_circuits_write_single_register() {
let store = Arc::new(SlaveStore::new());
let hook = WriteEchoHook::new();
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::WriteSingleRegister(100, 12345))
.await
.unwrap();
assert_eq!(rsp, Response::WriteSingleRegister(100, 12345));
assert_eq!(hooked.hook.last_written.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn before_call_short_circuits_write_multiple_registers() {
let store = Arc::new(SlaveStore::new());
let hook = WriteEchoHook::new();
let hooked = HookedService::new(store, hook);
let data = vec![10u16, 20, 30];
let rsp = hooked
.call(Request::WriteMultipleRegisters(
50,
std::borrow::Cow::Borrowed(&data),
))
.await
.unwrap();
assert_eq!(rsp, Response::WriteMultipleRegisters(50, 3));
assert_eq!(hooked.hook.last_written.lock().unwrap().len(), 3);
}
#[tokio::test]
async fn before_call_intercepts_diagnostic() {
struct DiagHook;
#[async_trait]
impl ServerHook for DiagHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::Diagnostic(sub_fn, data) = request {
Some(Response::Diagnostic(*sub_fn, *data))
} else {
None
}
}
}
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, DiagHook);
let rsp = hooked
.call(Request::Diagnostic(0x0000, 0xABCD))
.await
.unwrap();
assert_eq!(rsp, Response::Diagnostic(0x0000, 0xABCD));
}
struct SlaveAwareHook {
slave1_val: u16,
slave2_val: u16,
hits: [AtomicU8; 2],
}
#[async_trait]
impl ServerHook for SlaveAwareHook {
async fn before_call(&self, slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::ReadHoldingRegisters(_addr, qty) = request {
let val = match slave {
1 => {
self.hits[0].fetch_add(1, Ordering::Relaxed);
self.slave1_val
}
2 => {
self.hits[1].fetch_add(1, Ordering::Relaxed);
self.slave2_val
}
_ => return None,
};
Some(Response::ReadHoldingRegisters(vec![val; *qty as usize]))
} else {
None
}
}
}
#[tokio::test]
async fn before_call_routes_by_slave_id() {
let store = Arc::new(SlaveStore::new());
let hook = SlaveAwareHook {
slave1_val: 111,
slave2_val: 222,
hits: [AtomicU8::new(0), AtomicU8::new(0)],
};
let hooked = HookedService::new(store, hook);
let rsp1 = server::context::SLAVE_ID
.scope(1, async {
hooked.call(Request::ReadHoldingRegisters(0, 2)).await
})
.await
.unwrap();
assert_eq!(rsp1, Response::ReadHoldingRegisters(vec![111, 111]));
assert_eq!(hooked.hook.hits[0].load(Ordering::Relaxed), 1);
let rsp2 = server::context::SLAVE_ID
.scope(2, async {
hooked.call(Request::ReadHoldingRegisters(0, 2)).await
})
.await
.unwrap();
assert_eq!(rsp2, Response::ReadHoldingRegisters(vec![222, 222]));
assert_eq!(hooked.hook.hits[1].load(Ordering::Relaxed), 1);
let rsp3 = server::context::SLAVE_ID
.scope(3, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
assert_eq!(rsp3, Response::ReadHoldingRegisters(vec![0])); }
struct AddressRouterHook;
#[async_trait]
impl ServerHook for AddressRouterHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::ReadHoldingRegisters(addr, qty) if *addr < 100 => {
let regs: Vec<u16> = (*addr..*addr + *qty).collect();
Some(Response::ReadHoldingRegisters(regs))
}
Request::ReadHoldingRegisters(addr, qty) if *addr < 200 => {
let regs: Vec<u16> = (*addr..*addr + *qty).map(|a| a * 10).collect();
Some(Response::ReadHoldingRegisters(regs))
}
_ => None,
}
}
}
#[tokio::test]
async fn before_call_routes_by_address_range() {
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, AddressRouterHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![0, 1, 2]));
let rsp = hooked
.call(Request::ReadHoldingRegisters(100, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![1000, 1010, 1020]));
let rsp = hooked
.call(Request::ReadHoldingRegisters(500, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![0]));
}
struct ScaleHook {
factor: u16,
}
#[async_trait]
impl ServerHook for ScaleHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Ok(Response::ReadHoldingRegisters(regs)) => {
let scaled: Vec<u16> = regs.iter().map(|v| v.saturating_mul(self.factor)).collect();
Ok(Response::ReadHoldingRegisters(scaled))
}
other => other,
}
}
}
#[tokio::test]
async fn after_call_transforms_success_values() {
let store = Arc::new(SlaveStore::with_holding_registers(&[
(0, 5),
(1, 10),
(2, 100),
]));
let hooked = HookedService::new(store, ScaleHook { factor: 10 });
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![50, 100, 1000]));
}
#[tokio::test]
async fn after_call_scale_hook_only_affects_matched_fc() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 7)]));
let hooked = HookedService::new(store, ScaleHook { factor: 2 });
let rsp = hooked
.call(Request::WriteSingleRegister(0, 99))
.await
.unwrap();
assert_eq!(rsp, Response::WriteSingleRegister(0, 99));
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![198])); }
struct ExceptionUpgradeHook;
#[async_trait]
impl ServerHook for ExceptionUpgradeHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Err(_) => Err(Exception::ServerDeviceFailure),
ok => ok,
}
}
}
#[tokio::test]
async fn after_call_upgrades_exception_type() {
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, ExceptionUpgradeHook);
let err = hooked
.call(Request::Diagnostic(0x00FF, 0))
.await
.unwrap_err();
assert_eq!(err, Exception::ServerDeviceFailure);
}
struct ReadOnlyHook;
#[async_trait]
impl ServerHook for ReadOnlyHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Ok(rsp) => match rsp.function_code().value() {
5 | 6 | 15 | 16 | 22 | 23 => Err(Exception::IllegalFunction),
_ => Ok(rsp),
},
err => err,
}
}
}
#[tokio::test]
async fn after_call_escalates_write_to_exception() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
let hooked = HookedService::new(store, ReadOnlyHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
let err = hooked
.call(Request::WriteSingleRegister(0, 99))
.await
.unwrap_err();
assert_eq!(err, Exception::IllegalFunction);
}
struct PerSlaveCounter {
counts: [AtomicU32; 256],
}
impl PerSlaveCounter {
fn new() -> Self {
Self {
counts: std::array::from_fn(|_| AtomicU32::new(0)),
}
}
}
#[async_trait]
impl ServerHook for PerSlaveCounter {
async fn after_call(
&self,
slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
self.counts[slave as usize].fetch_add(1, Ordering::Relaxed);
result
}
}
#[tokio::test]
async fn after_call_counts_per_slave() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 1)]));
let hook = PerSlaveCounter::new();
let hooked = HookedService::new(store, hook);
for _ in 0..3 {
server::context::SLAVE_ID
.scope(1, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
}
for _ in 0..2 {
server::context::SLAVE_ID
.scope(5, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
}
assert_eq!(hooked.hook.counts[1].load(Ordering::Relaxed), 3);
assert_eq!(hooked.hook.counts[5].load(Ordering::Relaxed), 2);
assert_eq!(hooked.hook.counts[0].load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn triple_layer_hooks() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 10)]));
let log = LoggingHook::new();
let short = ShortCircuitHook {
fixed_value: 42,
call_count: AtomicU8::new(0),
};
let scale = ScaleHook { factor: 2 };
type L2 = HookedService<Arc<SlaveStore>, ScaleHook>;
type L1 = HookedService<L2, ShortCircuitHook>;
let hooked: HookedService<L1, LoggingHook> = HookedService::new(
HookedService::new(HookedService::new(store, scale), short),
log,
);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42, 42]));
assert_eq!(hooked.hook.before_count.load(Ordering::Relaxed), 1);
assert_eq!(hooked.hook.after_count.load(Ordering::Relaxed), 1);
}
struct FixedService {
value: u16,
}
#[async_trait]
impl Service for FixedService {
async fn call(&self, request: Request<'_>) -> Result<Response, Exception> {
match request {
Request::ReadHoldingRegisters(_addr, qty) => Ok(Response::ReadHoldingRegisters(vec![
self.value;
qty as usize
])),
Request::WriteSingleRegister(addr, val) => Ok(Response::WriteSingleRegister(addr, val)),
_ => Err(Exception::IllegalFunction),
}
}
}
#[tokio::test]
async fn hooked_custom_service_with_transform() {
let svc = FixedService { value: 100 };
let hooked = HookedService::new(svc, ScaleHook { factor: 3 });
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![300, 300]));
let rsp = hooked
.call(Request::WriteSingleRegister(0, 55))
.await
.unwrap();
assert_eq!(rsp, Response::WriteSingleRegister(0, 55));
}
#[tokio::test]
async fn hooked_custom_service_with_before_call_short_circuit() {
let svc = FixedService { value: 100 };
struct ReadOnlyAllHook;
#[async_trait]
impl ServerHook for ReadOnlyAllHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::ReadHoldingRegisters(_addr, qty) => {
Some(Response::ReadHoldingRegisters(vec![999; *qty as usize]))
}
_ => None,
}
}
}
let hooked = HookedService::new(svc, ReadOnlyAllHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![999, 999, 999]));
let err = hooked.call(Request::ReadCoils(0, 1)).await.unwrap_err();
assert_eq!(err, Exception::IllegalFunction);
}
#[tokio::test]
async fn hooks_can_share_state_via_arc() {
let shared_counter = Arc::new(AtomicU32::new(0));
struct SharedHook {
counter: Arc<AtomicU32>,
}
#[async_trait]
impl ServerHook for SharedHook {
async fn before_call(&self, _slave: u8, _request: &Request<'_>) -> Option<Response> {
self.counter.fetch_add(1, Ordering::Relaxed);
None
}
}
let hook1 = SharedHook {
counter: shared_counter.clone(),
};
let hook2 = SharedHook {
counter: shared_counter.clone(),
};
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 1)]));
let inner = HookedService::new(store, hook1);
let outer = HookedService::new(inner, hook2);
outer
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
outer
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(shared_counter.load(Ordering::Relaxed), 4);
}
struct AccessControlHook {
allowed_slaves: Vec<u8>,
blocked_count: AtomicU32,
}
#[async_trait]
impl ServerHook for AccessControlHook {
async fn before_call(&self, slave: u8, _request: &Request<'_>) -> Option<Response> {
if !self.allowed_slaves.contains(&slave) {
self.blocked_count.fetch_add(1, Ordering::Relaxed);
let fc = _request.function_code().value();
Some(Response::Exception(fc, Exception::IllegalFunction))
} else {
None
}
}
}
#[tokio::test]
async fn access_control_blocks_unauthorized_slaves() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 99)]));
let hook = AccessControlHook {
allowed_slaves: vec![1, 2],
blocked_count: AtomicU32::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = server::context::SLAVE_ID
.scope(1, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![99]));
let rsp = server::context::SLAVE_ID
.scope(3, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
match rsp {
Response::Exception(3, Exception::IllegalFunction) => {} other => panic!("expected Exception(3, IllegalFunction), got {other:?}"),
}
assert_eq!(hooked.hook.blocked_count.load(Ordering::Relaxed), 1);
}
struct LinearTransformHook {
offset: i32,
scale: i32,
}
#[async_trait]
impl ServerHook for LinearTransformHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Ok(Response::ReadHoldingRegisters(regs)) => {
let transformed: Vec<u16> = regs
.iter()
.map(|&v| {
let y = self.offset + self.scale * v as i32;
y.clamp(0, u16::MAX as i32) as u16
})
.collect();
Ok(Response::ReadHoldingRegisters(transformed))
}
other => other,
}
}
}
#[tokio::test]
async fn linear_transform_hook() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 100), (1, 200)]));
let hooked = HookedService::new(
store,
LinearTransformHook {
offset: 50,
scale: 2,
},
);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![250, 450]));
}
struct AuditHook {
log: Mutex<Vec<(u8, String, String)>>,
}
impl AuditHook {
fn new() -> Self {
Self {
log: Mutex::new(Vec::new()),
}
}
}
#[async_trait]
impl ServerHook for AuditHook {
async fn before_call(&self, slave: u8, request: &Request<'_>) -> Option<Response> {
self.log
.lock()
.unwrap()
.push((slave, format!("REQ {:?}", request), String::new()));
None
}
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
let entry = match &result {
Ok(rsp) => format!("OK {:?}", rsp),
Err(ex) => format!("ERR {:?}", ex),
};
if let Some(last) = self.log.lock().unwrap().last_mut() {
last.2 = entry;
}
result
}
}
#[tokio::test]
async fn audit_hook_records_request_response_pairs() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
let hooked = HookedService::new(store, AuditHook::new());
server::context::SLAVE_ID
.scope(1, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
server::context::SLAVE_ID
.scope(1, async {
hooked.call(Request::WriteSingleRegister(0, 99)).await
})
.await
.unwrap();
let log = hooked.hook.log.lock().unwrap();
assert_eq!(log.len(), 2);
assert!(
log[0].1.contains("ReadHoldingRegisters"),
"should log request"
);
assert!(
log[0].2.contains("ReadHoldingRegisters"),
"should log response"
);
assert!(
log[1].1.contains("WriteSingleRegister"),
"should log request"
);
assert!(
log[1].2.contains("WriteSingleRegister"),
"should log response"
);
}
struct RateLimitHook {
max_per_window: u32,
call_count: AtomicU32,
}
#[async_trait]
impl ServerHook for RateLimitHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
let fc = request.function_code().value();
let count = self.call_count.fetch_add(1, Ordering::Relaxed);
if count >= self.max_per_window {
Some(Response::Exception(fc, Exception::ServerDeviceFailure))
} else {
None
}
}
}
#[tokio::test]
async fn rate_limit_hook_blocks_after_threshold() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
let hook = RateLimitHook {
max_per_window: 2,
call_count: AtomicU32::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
match rsp {
Response::Exception(3, Exception::ServerDeviceFailure) => {}
other => panic!("expected rate-limit exception, got {other:?}"),
}
}
#[tokio::test]
async fn concurrent_calls_to_hooked_service() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 0)]));
let hook = LoggingHook::new();
let hooked = Arc::new(HookedService::new(store, hook));
let mut handles = Vec::new();
for i in 0..16 {
let svc = hooked.clone();
handles.push(tokio::spawn(async move {
let val = i as u16;
svc.call(Request::WriteSingleRegister(0, val))
.await
.unwrap();
let rsp = svc.call(Request::ReadHoldingRegisters(0, 1)).await.unwrap();
match rsp {
Response::ReadHoldingRegisters(regs) => {
assert_eq!(regs.len(), 1);
}
other => panic!("unexpected response: {other:?}"),
}
}));
}
for h in handles {
h.await.unwrap();
}
assert_eq!(hooked.hook.before_count.load(Ordering::Relaxed), 32);
assert_eq!(hooked.hook.after_count.load(Ordering::Relaxed), 32);
}
#[tokio::test]
async fn slave_id_defaults_to_zero_outside_scope() {
struct SlaveCaptureHook {
seen_slave: AtomicU8,
}
#[async_trait]
impl ServerHook for SlaveCaptureHook {
async fn before_call(&self, slave: u8, _request: &Request<'_>) -> Option<Response> {
self.seen_slave.store(slave, Ordering::Relaxed);
None
}
}
let store = Arc::new(SlaveStore::new());
let hook = SlaveCaptureHook {
seen_slave: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(hooked.hook.seen_slave.load(Ordering::Relaxed), 0xFF);
}
#[tokio::test]
async fn slave_id_scoped_to_specific_value() {
struct SlaveCaptureHook {
seen_slave: AtomicU8,
}
#[async_trait]
impl ServerHook for SlaveCaptureHook {
async fn before_call(&self, slave: u8, _request: &Request<'_>) -> Option<Response> {
self.seen_slave.store(slave, Ordering::Relaxed);
None
}
}
let store = Arc::new(SlaveStore::new());
let hook = SlaveCaptureHook {
seen_slave: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
server::context::SLAVE_ID
.scope(42, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
assert_eq!(hooked.hook.seen_slave.load(Ordering::Relaxed), 42);
server::context::SLAVE_ID
.scope(247, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
.unwrap();
assert_eq!(hooked.hook.seen_slave.load(Ordering::Relaxed), 247);
}
#[tokio::test]
async fn slave_id_nested_scope_uses_inner_value() {
struct SlaveCaptureHook {
seen_before: AtomicU8,
seen_after: AtomicU8,
}
#[async_trait]
impl ServerHook for SlaveCaptureHook {
async fn before_call(&self, slave: u8, _request: &Request<'_>) -> Option<Response> {
self.seen_before.store(slave, Ordering::Relaxed);
None
}
async fn after_call(
&self,
slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
self.seen_after.store(slave, Ordering::Relaxed);
result
}
}
let store = Arc::new(SlaveStore::new());
let hook = SlaveCaptureHook {
seen_before: AtomicU8::new(0),
seen_after: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
server::context::SLAVE_ID
.scope(99, async {
server::context::SLAVE_ID
.scope(77, async {
hooked.call(Request::ReadHoldingRegisters(0, 1)).await
})
.await
})
.await
.unwrap();
assert_eq!(hooked.hook.seen_before.load(Ordering::Relaxed), 77);
assert_eq!(hooked.hook.seen_after.load(Ordering::Relaxed), 77);
}
#[tokio::test]
async fn noop_hook_passthrough() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
let hooked = HookedService::new(store.clone(), NoopHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
let rsp = store
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42]));
hooked
.call(Request::WriteSingleRegister(0, 99))
.await
.unwrap();
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![99]));
}
#[tokio::test]
async fn before_call_short_circuit() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 10)]));
let hook = ShortCircuitHook {
fixed_value: 999,
call_count: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 3))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![999, 999, 999]));
let rsp = hooked
.call(Request::ReadHoldingRegisters(1, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![999, 999]));
}
#[tokio::test]
async fn before_call_pass_through_unmatched() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(5, 77)]));
let hook = ShortCircuitHook {
fixed_value: 111,
call_count: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
hooked
.call(Request::WriteSingleRegister(5, 88))
.await
.unwrap();
let rsp = hooked
.call(Request::ReadHoldingRegisters(5, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![111]));
}
#[tokio::test]
async fn after_call_suppresses_exception() {
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, SuppressExceptionHook);
let rsp = hooked.call(Request::Diagnostic(0x00FF, 0)).await.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![0]));
}
#[tokio::test]
async fn after_call_passthrough_success() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 55)]));
let hooked = HookedService::new(store, SuppressExceptionHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![55]));
}
#[tokio::test]
async fn layered_hooks() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 1)]));
let inner_hook = ShortCircuitHook {
fixed_value: 42,
call_count: AtomicU8::new(0),
};
let outer_hook = LoggingHook::new();
let hooked: HookedService<HookedService<Arc<SlaveStore>, ShortCircuitHook>, LoggingHook> =
HookedService::new(HookedService::new(store, inner_hook), outer_hook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42, 42]));
assert_eq!(hooked.hook.before_count.load(Ordering::Relaxed), 1);
assert_eq!(hooked.hook.after_count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn logging_hook_counts_calls() {
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 1), (1, 2)]));
let hook = LoggingHook::new();
let hooked = HookedService::new(store, hook);
hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
hooked
.call(Request::WriteSingleRegister(0, 99))
.await
.unwrap();
assert_eq!(hooked.hook.before_count.load(Ordering::Relaxed), 3);
assert_eq!(hooked.hook.after_count.load(Ordering::Relaxed), 3);
}
#[tokio::test]
async fn unknown_function_code_reaches_inner() {
let store = Arc::new(SlaveStore::new());
let hook = ShortCircuitHook {
fixed_value: 1,
call_count: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = hooked.call(Request::ReadCoils(0, 1)).await.unwrap();
match rsp {
Response::ReadCoils(bits) => assert_eq!(bits.len(), 1),
other => panic!("expected ReadCoils, got {other:?}"),
}
}
#[tokio::test]
async fn hooked_service_is_service() {
fn _assert_service(_: impl Service) {}
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, NoopHook);
_assert_service(hooked);
}
#[tokio::test]
async fn hook_with_only_before_call_implemented() {
struct BeforeOnlyHook;
#[async_trait]
impl ServerHook for BeforeOnlyHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::ReadHoldingRegisters(_, _) = request {
Some(Response::ReadHoldingRegisters(vec![77]))
} else {
None
}
}
}
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 10)]));
let hooked = HookedService::new(store, BeforeOnlyHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![77]));
let rsp = hooked.call(Request::ReadCoils(0, 1)).await.unwrap();
assert!(matches!(rsp, Response::ReadCoils(_)));
}
#[tokio::test]
async fn hook_with_only_after_call_implemented() {
struct AfterOnlyHook;
#[async_trait]
impl ServerHook for AfterOnlyHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Ok(Response::ReadHoldingRegisters(mut regs)) => {
regs.push(0xFFFF);
Ok(Response::ReadHoldingRegisters(regs))
}
other => other,
}
}
}
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 1), (1, 2)]));
let hooked = HookedService::new(store, AfterOnlyHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![1, 2, 0xFFFF]));
}
#[tokio::test]
async fn hook_handles_disconnect_request() {
struct DisconnectHook {
seen_disconnect: AtomicU8,
}
#[async_trait]
impl ServerHook for DisconnectHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if matches!(request, Request::Disconnect) {
self.seen_disconnect.fetch_add(1, Ordering::Relaxed);
}
None
}
}
let store = Arc::new(SlaveStore::new());
let hook = DisconnectHook {
seen_disconnect: AtomicU8::new(0),
};
let hooked = HookedService::new(store, hook);
let result = hooked.call(Request::Disconnect).await;
assert!(result.is_err());
assert_eq!(hooked.hook.seen_disconnect.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn before_call_intercepts_mask_write_register() {
struct MaskWriteHook {
last_and: Mutex<u16>,
last_or: Mutex<u16>,
}
#[async_trait]
impl ServerHook for MaskWriteHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::MaskWriteRegister(addr, and_mask, or_mask) = request {
*self.last_and.lock().unwrap() = *and_mask;
*self.last_or.lock().unwrap() = *or_mask;
Some(Response::MaskWriteRegister(*addr, *and_mask, *or_mask))
} else {
None
}
}
}
let store = Arc::new(SlaveStore::new());
let hook = MaskWriteHook {
last_and: Mutex::new(0),
last_or: Mutex::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::MaskWriteRegister(10, 0xFF00, 0x00FF))
.await
.unwrap();
assert_eq!(rsp, Response::MaskWriteRegister(10, 0xFF00, 0x00FF));
assert_eq!(*hooked.hook.last_and.lock().unwrap(), 0xFF00);
assert_eq!(*hooked.hook.last_or.lock().unwrap(), 0x00FF);
}
#[tokio::test]
async fn before_call_intercepts_read_write_multiple_registers() {
struct RwMultiHook;
#[async_trait]
impl ServerHook for RwMultiHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::ReadWriteMultipleRegisters(_rd_addr, rd_qty, _wr_addr, _wr_data) =
request
{
let regs = vec![0xAAAA; *rd_qty as usize];
Some(Response::ReadWriteMultipleRegisters(regs))
} else {
None
}
}
}
let store = Arc::new(SlaveStore::new());
let hooked = HookedService::new(store, RwMultiHook);
let rsp = hooked
.call(Request::ReadWriteMultipleRegisters(
0,
4,
100,
std::borrow::Cow::Borrowed(&[1u16, 2]),
))
.await
.unwrap();
assert_eq!(rsp, Response::ReadWriteMultipleRegisters(vec![0xAAAA; 4]));
}
#[tokio::test]
async fn before_call_can_return_exception_response() {
struct AlwaysFailHook;
#[async_trait]
impl ServerHook for AlwaysFailHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
let fc = request.function_code().value();
Some(Response::Exception(fc, Exception::ServerDeviceFailure))
}
}
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 42)]));
let hooked = HookedService::new(store, AlwaysFailHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 1))
.await
.unwrap();
match rsp {
Response::Exception(3, Exception::ServerDeviceFailure) => {}
other => panic!("expected Exception, got {other:?}"),
}
}
#[tokio::test]
async fn after_call_can_convert_exception_to_matched_response_type() {
struct GracefulDegradeHook;
#[async_trait]
impl ServerHook for GracefulDegradeHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match result {
Err(Exception::IllegalDataAddress) => Ok(Response::ReadHoldingRegisters(vec![])),
other => other,
}
}
}
struct AlwaysIllegalAddrService;
#[async_trait]
impl Service for AlwaysIllegalAddrService {
async fn call(&self, _request: Request<'_>) -> Result<Response, Exception> {
Err(Exception::IllegalDataAddress)
}
}
let svc = AlwaysIllegalAddrService;
let hooked = HookedService::new(svc, GracefulDegradeHook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(5000, 10))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![]));
}
#[tokio::test]
async fn after_call_not_called_when_before_call_short_circuits() {
let after_called = Arc::new(AtomicU8::new(0));
struct ShortBeforeHook;
#[async_trait]
impl ServerHook for ShortBeforeHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
if let Request::ReadHoldingRegisters(_addr, qty) = request {
Some(Response::ReadHoldingRegisters(vec![1; *qty as usize]))
} else {
None
}
}
}
struct AfterFlagHook {
flag: Arc<AtomicU8>,
}
#[async_trait]
impl ServerHook for AfterFlagHook {
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
self.flag.fetch_add(1, Ordering::Relaxed);
result
}
}
let store = Arc::new(SlaveStore::with_holding_registers(&[(0, 99)]));
let inner = HookedService::new(store, ShortBeforeHook);
let outer = HookedService::new(
inner,
AfterFlagHook {
flag: after_called.clone(),
},
);
let rsp = outer
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![1, 1]));
assert_eq!(after_called.load(Ordering::Relaxed), 1);
let rsp = outer.call(Request::ReadCoils(0, 1)).await.unwrap();
assert!(matches!(rsp, Response::ReadCoils(_)));
assert_eq!(after_called.load(Ordering::Relaxed), 2);
}
#[tokio::test]
async fn before_call_intercepts_multiple_fc_types_in_one_hook() {
struct MultiTypeHook {
read_val: u16,
last_coil_addr: Mutex<u16>,
}
#[async_trait]
impl ServerHook for MultiTypeHook {
async fn before_call(&self, _slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::ReadHoldingRegisters(_addr, qty) => {
Some(Response::ReadHoldingRegisters(vec![
self.read_val;
*qty as usize
]))
}
Request::ReadCoils(addr, qty) => {
*self.last_coil_addr.lock().unwrap() = *addr;
Some(Response::ReadCoils(vec![true; *qty as usize]))
}
Request::WriteSingleCoil(addr, val) => Some(Response::WriteSingleCoil(*addr, *val)),
_ => None,
}
}
}
let store = Arc::new(SlaveStore::new());
let hook = MultiTypeHook {
read_val: 42,
last_coil_addr: Mutex::new(0),
};
let hooked = HookedService::new(store, hook);
let rsp = hooked
.call(Request::ReadHoldingRegisters(0, 2))
.await
.unwrap();
assert_eq!(rsp, Response::ReadHoldingRegisters(vec![42, 42]));
let rsp = hooked.call(Request::ReadCoils(100, 3)).await.unwrap();
assert_eq!(rsp, Response::ReadCoils(vec![true; 3]));
assert_eq!(*hooked.hook.last_coil_addr.lock().unwrap(), 100);
let rsp = hooked
.call(Request::WriteSingleCoil(50, true))
.await
.unwrap();
assert_eq!(rsp, Response::WriteSingleCoil(50, true));
let rsp = hooked
.call(Request::ReadDiscreteInputs(0, 1))
.await
.unwrap();
assert!(matches!(rsp, Response::ReadDiscreteInputs(_)));
}