use crate::limits::MAX_DNS_MESSAGE_LEN;
use crate::Message;
use std::net::SocketAddr;
pub struct AsyncClient {
server: SocketAddr,
socket: Option<tokio::net::UdpSocket>,
}
impl AsyncClient {
pub fn new(server: SocketAddr) -> Self {
Self {
server,
socket: None,
}
}
pub async fn exchange(&mut self, query: &Message) -> Result<Message, crate::Error> {
if self.socket.is_none() {
log::trace!("async UDP target={}", self.server);
let socket = tokio::net::UdpSocket::bind("0.0.0.0:0").await?;
socket.connect(self.server).await?;
log::trace!(
"async UDP connected local={} peer={}",
socket.local_addr()?,
socket.peer_addr()?
);
self.socket = Some(socket);
} else {
log::trace!("async UDP reusing connected socket peer={}", self.server);
}
let result: std::io::Result<Message> = async {
let socket = self.socket.as_mut().ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::NotConnected, "UDP socket unavailable")
})?;
let request = query.to_vec()?;
log::trace!(
"async UDP sending {} bytes to {}",
request.len(),
self.server
);
socket.send(&request).await?;
let mut response = [0; MAX_DNS_MESSAGE_LEN];
let length = socket.recv(&mut response).await?;
log::trace!("async UDP received {length} bytes from {}", self.server);
Message::from_slice(&response[..length])
}
.await;
match result {
Ok(response) => Ok(response),
Err(error) => {
log::trace!("async UDP discarding socket after error: {error}");
self.socket = None;
Err(error.into())
}
}
}
}
#[cfg(test)]
mod tests {
use super::AsyncClient;
use crate::Message;
use std::net::SocketAddr;
use tokio::net::UdpSocket;
#[tokio::test]
async fn exchanges_with_local_server() {
let server_socket = UdpSocket::bind("127.0.0.1:0")
.await
.expect("bind test socket");
let address: SocketAddr = server_socket
.local_addr()
.expect("read test socket address");
let server = tokio::spawn(async move {
let mut request = [0; 512];
let (_, peer) = server_socket
.recv_from(&mut request)
.await
.expect("read request");
server_socket
.send_to(&[0; 12], peer)
.await
.expect("write response");
});
let mut client = AsyncClient::new(address);
assert!(client.exchange(&Message::default()).await.is_ok());
server.await.expect("join test server");
}
}