diff --git a/Cargo.lock b/Cargo.lock index 898c7bd..6d71e2a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -216,10 +216,13 @@ dependencies = [ name = "crete" version = "0.1.2" dependencies = [ + "byteorder", "clap", "clap_complete", "gix-discover", "minijinja", + "netascii", + "num_enum", "optional_struct", "reqwest", "schemars", @@ -228,6 +231,7 @@ dependencies = [ "serde_json", "serde_yaml", "thiserror", + "tokio", "walkdir", ] @@ -447,7 +451,7 @@ dependencies = [ "bstr", "gix-date", "gix-error", - "winnow", + "winnow 0.7.15", ] [[package]] @@ -567,7 +571,7 @@ dependencies = [ "itoa", "smallvec", "thiserror", - "winnow", + "winnow 0.7.15", ] [[package]] @@ -600,7 +604,7 @@ dependencies = [ "gix-validate", "memmap2", "thiserror", - "winnow", + "winnow 0.7.15", ] [[package]] @@ -1125,6 +1129,34 @@ dependencies = [ "tempfile", ] +[[package]] +name = "netascii" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "205c45083b6725eb4671392c7c399510ad1498dd0275f1a2e68a70a34ebdf1b1" + +[[package]] +name = "num_enum" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d0bca838442ec211fa11de3a8b0e0e8f3a4522575b5c4c06ed722e005036f26" +dependencies = [ + "num_enum_derive", + "rustversion", +] + +[[package]] +name = "num_enum_derive" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "680998035259dcfcafe653688bf2aa6d3e2dc05e98be6ab46afb089dc84f1df8" +dependencies = [ + "proc-macro-crate", + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -1277,6 +1309,15 @@ dependencies = [ "syn", ] +[[package]] +name = "proc-macro-crate" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f" +dependencies = [ + "toml_edit", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -1628,6 +1669,16 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + [[package]] name = "slab" version = "0.4.12" @@ -1753,18 +1804,32 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.50.0" +version = "1.52.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "27ad5e34374e03cfffefc301becb44e9dc3c17584f414349ebe29ed26661822d" +checksum = "b67dee974fe86fd92cc45b7a95fdd2f99a36a6d7b0d431a231178d3d670bbcc6" dependencies = [ "bytes", "libc", "mio", + "parking_lot", "pin-project-lite", + "signal-hook-registry", "socket2", + "tokio-macros", "windows-sys", ] +[[package]] +name = "tokio-macros" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "tokio-native-tls" version = "0.3.1" @@ -1788,6 +1853,36 @@ dependencies = [ "tokio", ] +[[package]] +name = "toml_datetime" +version = "1.1.1+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7" +dependencies = [ + "serde_core", +] + +[[package]] +name = "toml_edit" +version = "0.25.11+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b59c4d22ed448339746c59b905d24568fcbb3ab65a500494f7b8c3e97739f2b" +dependencies = [ + "indexmap", + "toml_datetime", + "toml_parser", + "winnow 1.0.2", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow 1.0.2", +] + [[package]] name = "tower" version = "0.5.3" @@ -2115,6 +2210,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "winnow" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee1708bef14716a11bae175f579062d4554d95be2c6829f518df847b7b3fdd0" +dependencies = [ + "memchr", +] + [[package]] name = "wit-bindgen" version = "0.51.0" diff --git a/Cargo.toml b/Cargo.toml index 616d8d8..f20d0c6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,10 +14,13 @@ schemars = { version = "1.2.1", features = ["semver1"] } serde_json = "1.0.149" [dependencies] +byteorder = "1.5.0" clap = { version = "4.6.0", features = ["derive"] } clap_complete = "4.6.0" gix-discover = { version = "0.49.0", features = ["sha1"] } minijinja = { version = "2.18.0", features = ["json", "loader"] } +netascii = "0.1.0" +num_enum = "0.7.6" optional_struct = "0.5.2" reqwest = { version = "0.13.2", default-features = false, features = [ "blocking", @@ -30,6 +33,7 @@ serde = { version = "1.0.228", features = ["derive"] } serde_json = { workspace = true } serde_yaml = "0.9.34" thiserror = "2.0.18" +tokio = { version = "1.52.1", features = ["full"] } walkdir = "2.5.0" [lib] diff --git a/src/lib.rs b/src/lib.rs index 8b12a92..7685d78 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -8,6 +8,7 @@ pub mod node; pub mod patch; pub mod schematic; pub mod secret; +pub mod tftp; pub(crate) static REPO_PATH: OnceLock = OnceLock::new(); diff --git a/src/tftp/error.rs b/src/tftp/error.rs new file mode 100644 index 0000000..8fafc2d --- /dev/null +++ b/src/tftp/error.rs @@ -0,0 +1,20 @@ +use std::io; +use std::string::FromUtf8Error; + +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum Error { + #[error("Invalid op code '{0}'")] + InvalidOpCode(u16), + #[error("Invalid error code '{0}'")] + InvalidErrorCode(u16), + #[error("Invalid mode '{0}'")] + InvalidMode(String), + #[error("String contains invalid characters")] + InvalidString(#[from] FromUtf8Error), + #[error("Unterminated string")] + UnterminatedString, + #[error("Io error: {0}")] + Io(#[from] io::Error), +} diff --git a/src/tftp/mod.rs b/src/tftp/mod.rs new file mode 100644 index 0000000..8eab9f3 --- /dev/null +++ b/src/tftp/mod.rs @@ -0,0 +1,7 @@ +mod error; +mod packet; +mod server; + +pub use error::Error; +pub use packet::Packet; +pub use server::{Error as ServerError, serve}; diff --git a/src/tftp/packet.rs b/src/tftp/packet.rs new file mode 100644 index 0000000..28612a3 --- /dev/null +++ b/src/tftp/packet.rs @@ -0,0 +1,353 @@ +use std::io::{BufRead, BufReader, BufWriter, Cursor, Read, Write}; +use std::net::SocketAddr; + +use byteorder::{NetworkEndian, ReadBytesExt, WriteBytesExt}; +use num_enum::{IntoPrimitive, TryFromPrimitive}; +use tokio::net::UdpSocket; + +use crate::tftp::Error; + +// 2 bytes op code + 2 bytes block number + max 512 bytes data +pub(crate) const PACKET_SIZE: usize = 2 + 2 + 512; + +fn decode_null_terminated(reader: &mut BufReader) -> Result +where + R: Read + ?Sized, +{ + let mut text = Vec::new(); + reader.read_until(0, &mut text)?; + + // Make sure we ended with a null terminator + match text.pop() { + Some(0) => {} + None | Some(_) => return Err(Error::UnterminatedString), + } + + Ok(String::from_utf8(text)?) +} + +fn encode_null_terminated(writer: &mut BufWriter, text: impl AsRef) -> Result<(), Error> +where + W: Write + ?Sized, +{ + let text = text.as_ref().as_bytes(); + writer.write_all(text)?; + + // Add null terminator + writer.write_u8(0)?; + + Ok(()) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Mode { + NetAscii, + Octet, + Mail, +} + +impl Mode { + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let mut mode = decode_null_terminated(reader)?; + mode.make_ascii_lowercase(); + let mode = mode; + + Ok(match mode.as_str() { + "netascii" => Self::NetAscii, + "octet" => Self::Octet, + "mail" => Self::Mail, + _ => return Err(Error::InvalidMode(mode)), + }) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + let mode = match &self { + Mode::NetAscii => "netascii", + Mode::Octet => "octet", + Mode::Mail => "mail", + }; + + encode_null_terminated(writer, mode)?; + + Ok(()) + } +} + +#[derive(Debug, PartialEq, Eq)] +pub struct Request { + filename: String, + mode: Mode, +} + +impl Request { + pub fn filename(&self) -> &str { + &self.filename + } + + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let filename = decode_null_terminated(reader)?; + let mode = Mode::decode(reader)?; + + Ok(Self { filename, mode }) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + encode_null_terminated(writer, &self.filename)?; + self.mode.encode(writer)?; + + Ok(()) + } + + pub fn mode(&self) -> Mode { + self.mode + } +} + +#[derive(Debug)] +pub struct Data { + block: u16, + data: Vec, +} + +impl Data { + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let block = reader.read_u16::()?; + + let mut data = Vec::new(); + reader.read_to_end(&mut data)?; + + Ok(Self { data, block }) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + writer.write_u16::(self.block)?; + writer.write_all(&self.data)?; + + Ok(()) + } +} + +#[derive(Debug)] +pub struct Ack(u16); + +impl Ack { + pub fn block(&self) -> u16 { + self.0 + } + + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let number = reader.read_u16::()?; + + Ok(Self(number)) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + writer.write_u16::(self.0)?; + + Ok(()) + } +} + +#[derive(Debug, TryFromPrimitive, IntoPrimitive, Clone, Copy)] +#[repr(u16)] +pub enum ErrorCode { + NotDefined, + FileNotFound, + AccessViolation, + DiskFull, + Illegal, + UnknownTransferId, + FileAlreadyExists, + NoSuchUser, +} + +impl ErrorCode { + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let code = reader.read_u16::()?; + Self::try_from_primitive(code).map_err(|err| Error::InvalidErrorCode(err.number)) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + let code: u16 = (*self).into(); + + writer.write_u16::(code)?; + + Ok(()) + } +} + +#[derive(Debug)] +pub struct ErrorMessage { + code: ErrorCode, + message: String, +} + +impl ErrorMessage { + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let code = ErrorCode::decode(reader)?; + let message = decode_null_terminated(reader)?; + + Ok(Self { code, message }) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + self.code.encode(writer)?; + + encode_null_terminated(writer, &self.message)?; + + Ok(()) + } +} + +#[derive(Debug)] +pub enum Packet { + Read(Request), + Write(Request), + Data(Data), + Ack(Ack), + Error(ErrorMessage), +} + +impl Packet { + pub fn data(block: u16, data: &[u8]) -> Self { + Self::Data(Data { + block, + data: data.into(), + }) + } + + pub fn ack(block: u16) -> Self { + Self::Ack(Ack(block)) + } + + pub fn error(code: ErrorCode, message: impl Into) -> Self { + Self::Error(ErrorMessage { + code, + message: message.into(), + }) + } + + fn decode(reader: &mut BufReader) -> Result + where + R: Read + ?Sized, + { + let op_code = reader.read_u16::()?; + + Ok(match op_code { + 1 => Self::Read(Request::decode(reader)?), + 2 => Self::Write(Request::decode(reader)?), + 3 => Self::Data(Data::decode(reader)?), + 4 => Self::Ack(Ack::decode(reader)?), + 5 => Self::Error(ErrorMessage::decode(reader)?), + _ => return Err(Error::InvalidOpCode(op_code)), + }) + } + + fn encode(&self, writer: &mut BufWriter) -> Result<(), Error> + where + W: Write + ?Sized, + { + match &self { + Self::Read(inner) => { + writer.write_u16::(1)?; + inner.encode(writer)?; + } + Self::Write(inner) => { + writer.write_u16::(2)?; + inner.encode(writer)?; + } + Self::Data(inner) => { + writer.write_u16::(3)?; + inner.encode(writer)?; + } + Self::Ack(inner) => { + writer.write_u16::(4)?; + inner.encode(writer)?; + } + Self::Error(inner) => { + writer.write_u16::(5)?; + inner.encode(writer)?; + } + }; + + Ok(()) + } + + pub async fn recv_from(socket: &UdpSocket) -> Result<(Self, SocketAddr), Error> { + let mut buf = Vec::with_capacity(PACKET_SIZE); + let (_, src) = socket.recv_buf_from(&mut buf).await?; + + let packet = buf.try_into()?; + + Ok((packet, src)) + } + + pub async fn send_to(&self, socket: &UdpSocket, src: &SocketAddr) -> Result<(), Error> { + let buf: Vec<_> = self.try_into()?; + socket.send_to(&buf, src).await?; + + Ok(()) + } +} + +impl TryFrom<&Packet> for Vec { + type Error = Error; + + fn try_from(packet: &Packet) -> Result { + let mut buf = Vec::with_capacity(PACKET_SIZE); + let mut writer = BufWriter::new(Cursor::new(&mut buf)); + packet.encode(&mut writer)?; + + // Normally the writer would get dropped after we return, but since it has a mutable + // borrow in buf we have to drop it first. + drop(writer); + + Ok(buf) + } +} + +impl TryFrom> for Packet { + type Error = Error; + + fn try_from(buf: Vec) -> Result>>::Error> { + let mut reader = BufReader::new(Cursor::new(buf)); + + Packet::decode(&mut reader) + } +} diff --git a/src/tftp/server.rs b/src/tftp/server.rs new file mode 100644 index 0000000..1586c09 --- /dev/null +++ b/src/tftp/server.rs @@ -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>) + 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}") + } + } + }); + } +}