use std::io; use std::net::SocketAddr; use std::time::Duration; use thiserror::Error; use tokio::net::UdpSocket; use crate::tftp::packet::{ErrorCode, Mode}; use crate::tftp::{self, Packet}; #[derive(Debug, Error)] pub enum Error { #[error("{0}")] Tftp(#[from] tftp::Error), #[error("Io error: {0}")] Io(#[from] io::Error), #[error("Unsupported operation: {0:#?}")] UnsupportedOperation(Packet), #[error("File not found: {0}")] FileNotFound(String), #[error("Unsupported mode: {0:?}")] UnsupportedMode(Mode), #[error("Source changed: {0}")] SourceChanged(SocketAddr), #[error("UnexpectedPacket: {0:#?}")] UnexpectedPacket(Packet), #[error("Timeout")] Timeout, } pub async fn handle_connection( src: SocketAddr, packet: Packet, callback: impl (Fn(&str) -> Option>) + Send + Clone + 'static, ) -> Result<(), Error> { let socket = UdpSocket::bind("127.0.0.1:0").await?; println!( "{src}: Connecting from {}", socket.local_addr().expect("Socket is bound") ); let Packet::Read(request) = packet else { Packet::error(ErrorCode::NotDefined, "Operation not supported") .send_to(&socket, &src) .await?; return Err(Error::UnsupportedOperation(packet)); }; let filename = request.filename(); println!("{src}: {filename} [{:?}]", request.mode()); let Some(data) = callback(filename) else { Packet::error(ErrorCode::FileNotFound, "Unknown file") .send_to(&socket, &src) .await?; return Err(Error::FileNotFound(filename.into())); }; let mode = request.mode(); let data = match mode { // In case of netascii mode we need to re-encode to escape certain sequences Mode::NetAscii => netascii::Netascii::from_bytes(data).collect(), Mode::Octet => data, _ => { Packet::error( ErrorCode::NotDefined, format!("Mode {mode:?} is not supported"), ) .send_to(&socket, &src) .await?; return Err(Error::UnsupportedMode(mode)); } }; let mut chunks = data.chunks(512).enumerate(); let mut chunk = chunks.next(); let mut timeout = 0; loop { // Send the first byte if let Some((index, chunk)) = chunk { Packet::data((index + 1) as u16, chunk) .send_to(&socket, &src) .await?; } tokio::select! { recv = Packet::recv_from(&socket) => { timeout = 0; let (packet, new_src) = recv?; // I think this is might technically be possible and fine, but I don't // really know how to test it properly so we don't support it. if new_src != src { Packet::error(ErrorCode::Illegal, "Source changed") .send_to(&socket, &src) .await?; return Err(Error::SourceChanged(new_src)); } match packet { Packet::Ack(ack) => { // If we receive an ack for the current chunk go to the next one if let Some((index, _)) = chunk && (index + 1) as u16 == ack.block() { chunk = chunks.next(); } if chunk.is_none() { break; } } _ => { Packet::error(ErrorCode::Illegal, "Unexpected packet") .send_to(&socket, &src) .await?; return Err(Error::UnexpectedPacket(packet)); } } }, _ = tokio::time::sleep(Duration::from_secs(1)) => { timeout += 1; if timeout == 5 { return Err(Error::Timeout); } } } } println!("{src}: Transmission complete"); Ok(()) } pub async fn serve( callback: impl (Fn(&str) -> Option>) + Send + Clone + 'static, ) -> Result<(), Error> { let socket = UdpSocket::bind("127.0.0.1:69").await?; println!("Server started"); loop { let (packet, src) = Packet::recv_from(&socket).await?; let callback = callback.clone(); tokio::spawn({ async move { if let Err(err) = handle_connection(src, packet, callback).await { println!("{src}: {err}") } } }); } }