use std::io;
use std::net::SocketAddr;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::UdpSocket;
use tokio::time::{Interval, MissedTickBehavior};
use super::wg::WgTunnel;
const RX_LEN: usize = 65_536;
const TIMER_PERIOD: Duration = Duration::from_millis(250);
pub struct WgDevice {
socket: UdpSocket,
tunnel: WgTunnel,
peer: Option<SocketAddr>,
rx: Box<[u8]>,
keepalive: Interval,
}
impl WgDevice {
#[must_use]
pub fn new(socket: UdpSocket, tunnel: WgTunnel) -> Self {
let mut keepalive = tokio::time::interval(TIMER_PERIOD);
keepalive.set_missed_tick_behavior(MissedTickBehavior::Delay);
Self {
socket,
tunnel,
peer: None,
rx: vec![0u8; RX_LEN].into_boxed_slice(),
keepalive,
}
}
}
impl AsyncRead for WgDevice {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let WgDevice {
socket,
tunnel,
peer,
rx,
keepalive,
} = self.get_mut();
loop {
while keepalive.poll_tick(cx).is_ready() {
if let Some(p) = *peer {
tunnel.tick(|b| {
let _ = socket.try_send_to(b, p);
});
}
}
let mut rb = ReadBuf::new(&mut rx[..]);
match socket.poll_recv_from(cx, &mut rb) {
Poll::Ready(Ok(addr)) => {
*peer = Some(addr);
let datagram = rb.filled();
let mut produced = false;
tunnel.decapsulate(
datagram,
|b| {
let _ = socket.try_send_to(b, addr);
},
|b| {
if !produced && b.len() <= buf.remaining() {
buf.put_slice(b);
produced = true;
}
},
);
if produced {
return Poll::Ready(Ok(()));
}
}
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Pending => return Poll::Pending,
}
}
}
}
impl AsyncWrite for WgDevice {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
data: &[u8],
) -> Poll<io::Result<usize>> {
let WgDevice {
socket,
tunnel,
peer,
..
} = self.get_mut();
if let Some(p) = *peer {
tunnel.encapsulate(data, |b| {
let _ = socket.try_send_to(b, p);
});
}
Poll::Ready(Ok(data.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}