This commit is contained in:
@@ -0,0 +1,162 @@
|
||||
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<Vec<u8>>) + 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<Vec<u8>>) + 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}")
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user