use std::io;
use std::net::SocketAddr;
use matter_commissioning::driver::AsyncDatagram;
use tokio::sync::{mpsc, Mutex};
pub(crate) struct HandshakeOutbound {
pub node_id: u64,
pub bytes: Vec<u8>,
pub peer: SocketAddr,
}
pub(crate) struct HandshakeSocket {
node_id: u64,
outbound: mpsc::Sender<HandshakeOutbound>,
inbound: Mutex<mpsc::Receiver<(Vec<u8>, SocketAddr)>>,
}
impl HandshakeSocket {
pub(crate) fn new(
node_id: u64,
outbound: mpsc::Sender<HandshakeOutbound>,
inbound: mpsc::Receiver<(Vec<u8>, SocketAddr)>,
) -> Self {
Self {
node_id,
outbound,
inbound: Mutex::new(inbound),
}
}
}
impl AsyncDatagram for HandshakeSocket {
async fn send_to(&self, buf: &[u8], peer: SocketAddr) -> io::Result<()> {
self.outbound
.send(HandshakeOutbound {
node_id: self.node_id,
bytes: buf.to_vec(),
peer,
})
.await
.map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "actor loop gone"))
}
async fn recv_from(&self) -> io::Result<(Vec<u8>, SocketAddr)> {
let mut rx = self.inbound.lock().await;
rx.recv()
.await
.ok_or_else(|| io::Error::new(io::ErrorKind::BrokenPipe, "actor loop gone"))
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)] mod tests {
use super::*;
#[tokio::test]
async fn bridges_outbound_and_inbound_over_channels() {
let (out_tx, mut out_rx) = mpsc::channel::<HandshakeOutbound>(4);
let (in_tx, in_rx) = mpsc::channel::<(Vec<u8>, SocketAddr)>(4);
let sock = HandshakeSocket::new(0x1234, out_tx, in_rx);
let peer: SocketAddr = "127.0.0.1:5540".parse().unwrap();
sock.send_to(b"sigma1", peer).await.unwrap();
let out = out_rx.recv().await.unwrap();
assert_eq!(out.node_id, 0x1234);
assert_eq!(out.bytes, b"sigma1");
assert_eq!(out.peer, peer);
in_tx.send((b"sigma2".to_vec(), peer)).await.unwrap();
let (bytes, from) = sock.recv_from().await.unwrap();
assert_eq!(bytes, b"sigma2");
assert_eq!(from, peer);
}
#[tokio::test]
async fn reports_broken_pipe_when_actor_gone() {
let (out_tx, out_rx) = mpsc::channel::<HandshakeOutbound>(1);
let (in_tx, in_rx) = mpsc::channel::<(Vec<u8>, SocketAddr)>(1);
let sock = HandshakeSocket::new(1, out_tx, in_rx);
let peer: SocketAddr = "127.0.0.1:5540".parse().unwrap();
drop(out_rx); let err = sock.send_to(b"x", peer).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::BrokenPipe);
drop(in_tx); let err = sock.recv_from().await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::BrokenPipe);
}
}