use super::{DatagramConfig, DatagramProtocol};
use crate::common::EchoServerTrait;
use crate::{EchoError, Result};
use async_trait::async_trait;
use std::sync::Arc;
use tokio::{signal, time::timeout};
use tracing::{error, info, warn};
pub struct DatagramEchoServer<P: DatagramProtocol> {
config: DatagramConfig,
protocol: std::marker::PhantomData<P>,
shutdown_signal: Arc<tokio::sync::broadcast::Sender<()>>,
}
impl<P: DatagramProtocol> DatagramEchoServer<P>
where
P::Error: Into<EchoError> + std::fmt::Display,
{
pub fn new(config: DatagramConfig) -> Self {
let (shutdown_signal, _) = tokio::sync::broadcast::channel(1);
Self {
config,
protocol: std::marker::PhantomData,
shutdown_signal: Arc::new(shutdown_signal),
}
}
}
#[async_trait]
impl<P: DatagramProtocol + Sync> EchoServerTrait for DatagramEchoServer<P>
where
P::Error: Into<EchoError> + std::fmt::Display,
{
async fn run(&self) -> Result<()> {
let socket = P::bind(&self.config).await.map_err(|e| e.into())?;
info!(address = %self.config.bind_addr, "Datagram echo server listening");
let mut buffer = vec![0; self.config.buffer_size];
let mut shutdown_rx = self.shutdown_signal.subscribe();
loop {
tokio::select! {
recv_result = timeout(self.config.read_timeout, P::recv_from(&socket, &mut buffer)) => {
match recv_result {
Ok(Ok((n, addr))) => {
let preview = String::from_utf8_lossy(&buffer[..n]);
info!(%addr, size = n, preview = %preview, "Received datagram");
if let Err(e) = P::send_to(&socket, &buffer[..n], addr).await {
error!(%addr, error = %e, "Failed to send echo response");
} else {
info!(%addr, size = n, "Echoed datagram");
}
}
Ok(Err(e)) => {
error!(error = %e, "Failed to receive datagram");
}
Err(_) => {
warn!("Receive timeout");
}
}
}
_ = signal::ctrl_c() => {
info!("Received shutdown signal, stopping server");
break;
}
_ = shutdown_rx.recv() => {
info!("Received internal shutdown signal, stopping server");
break;
}
}
}
info!("Datagram echo server stopped");
Ok(())
}
fn shutdown_signal(&self) -> tokio::sync::broadcast::Sender<()> {
self.shutdown_signal.as_ref().clone()
}
}