use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use oms_modbus::*;
#[derive(Clone)]
struct LoggingHook {
call_count: Arc<AtomicU64>,
}
#[async_trait]
impl ServerHook for LoggingHook {
async fn before_call(&self, slave: u8, request: &Request<'_>) -> Option<Response> {
let n = self.call_count.fetch_add(1, Ordering::Relaxed);
println!(" [{n:04}] ← slave={slave} {:?}", request);
None }
async fn after_call(
&self,
_slave: u8,
result: Result<Response, Exception>,
) -> Result<Response, Exception> {
match &result {
Ok(rsp) => println!(" → {:?}", rsp),
Err(ex) => println!(" → Exception: {:?}", ex),
}
result
}
}
#[derive(Clone)]
struct AccessControlHook {
admin_slaves: Vec<u8>,
}
#[async_trait]
impl ServerHook for AccessControlHook {
async fn before_call(&self, slave: u8, request: &Request<'_>) -> Option<Response> {
let is_write = matches!(
request,
Request::WriteSingleRegister(..)
| Request::WriteMultipleRegisters(..)
| Request::WriteSingleCoil(..)
| Request::WriteMultipleCoils(..)
| Request::MaskWriteRegister(..)
| Request::ReadWriteMultipleRegisters(..)
);
if is_write && !self.admin_slaves.contains(&slave) {
let fc = request.function_code().value();
println!(" 🚫 Access denied: slave={slave} attempted write FC={fc}");
return Some(Response::Exception(fc, Exception::IllegalFunction));
}
None
}
}
#[derive(Clone)]
struct ComputeHook {
sensor_value: Arc<std::sync::Mutex<u16>>,
}
impl ComputeHook {
fn new(initial: u16) -> Self {
Self {
sensor_value: Arc::new(std::sync::Mutex::new(initial)),
}
}
}
#[async_trait]
impl ServerHook for ComputeHook {
async fn before_call(&self, slave: u8, request: &Request<'_>) -> Option<Response> {
match request {
Request::ReadHoldingRegisters(addr, qty) if *addr >= 100 && *addr < 110 => {
let mut val = self.sensor_value.lock().unwrap();
*val = val.wrapping_add(1);
let regs: Vec<u16> = (0..*qty).map(|i| (*val).wrapping_add(i)).collect();
Some(Response::ReadHoldingRegisters(regs))
}
Request::ReadHoldingRegisters(addr, qty) if *addr == 200 => {
let base: u16 = 0x4F4D; let mut regs = vec![base, 0x5300 + slave as u16]; regs.resize(*qty as usize, 0);
Some(Response::ReadHoldingRegisters(regs))
}
_ => None,
}
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("═══ OMS Modbus — ServerHook Demo ═══\n");
println!("This example demonstrates three hook patterns:\n");
println!(" ① Logging — prints every request/response");
println!(" ② AccessCtrl — blocks writes from non-admin slaves");
println!(" ③ Compute — dynamic register values (Lua-like)\n");
let addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0);
let server = tcp::TcpServer::bind(addr).await?;
let bind_addr = server.local_addr()?;
let store = Arc::new(SlaveStore::with_holding_registers(&[
(0, 1234), (1, 5678), (2, 100), (3, 42), ]));
let logging = LoggingHook {
call_count: Arc::new(AtomicU64::new(0)),
};
let access = AccessControlHook {
admin_slaves: vec![1],
}; let compute = ComputeHook::new(1000);
type L2 = HookedService<Arc<SlaveStore>, ComputeHook>;
type L1 = HookedService<L2, AccessControlHook>;
let hooked: HookedService<L1, LoggingHook> = HookedService::new(
HookedService::new(HookedService::new(store, compute), access),
logging,
);
tokio::spawn(async move {
server.serve_forever(hooked).await.ok();
});
tokio::time::sleep(Duration::from_millis(50)).await;
println!("Server listening on {bind_addr}\n");
let client = tcp::TcpClient::connect_with_timeout(bind_addr, Duration::from_secs(3)).await?;
println!("─── Demo ①: Logging (watch the [NNNN] log lines) ───");
let regs = client.read_holding_registers(1, 0, 3).await?;
println!(" Result: {regs:?}\n");
println!("─── Demo ②: Dynamic Compute (address 100..109) ───");
for _ in 0..3 {
let regs = client.read_holding_registers(1, 100, 2).await?;
println!(" Computed[100..101] = {regs:?} ← increments each read");
}
println!();
println!("─── Demo ③: Device Info (address 200) ───");
let info = client.read_holding_registers(1, 200, 4).await?;
println!(" DeviceInfo = {info:?}\n");
println!("─── Demo ④: Write as admin slave 1 (allowed) ───");
client.write_single_register(1, 0, 9999).await?;
let regs = client.read_holding_registers(1, 0, 1).await?;
println!(" Holding[0] = {regs:?} ← write succeeded\n");
println!("─── Demo ⑤: Write as slave 2 (BLOCKED) ───");
match client.write_single_register(2, 0, 5555).await {
Ok(()) => println!(" Unexpected: write succeeded"),
Err(e) => println!(" ✗ {e} ← correctly blocked by AccessControlHook\n"),
}
println!("─── Demo ⑥: Read as slave 2 (allowed — reads not blocked) ───");
match client.read_holding_registers(2, 0, 1).await {
Ok(regs) => println!(" Holding[0] = {regs:?}\n"),
Err(e) => println!(" ✗ {e}\n"),
}
println!("═══ Summary ═══");
println!(" ✓ LoggingHook — recorded every request/response (see [NNNN] above)");
println!(" ✓ ComputeHook — short-circuited reads to addresses 100-109, 200");
println!(" ✓ AccessCtrlHook — blocked write from slave 2, allowed slave 1");
println!();
println!("Integrate your own hook: implement ServerHook, wrap with HookedService.");
println!("Real-world examples: Lua script execution, database lookup, MQTT bridge.");
Ok(())
}