diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/configuration.rs | 44 | ||||
| -rw-r--r-- | src/error.rs | 20 | ||||
| -rw-r--r-- | src/handshake.rs | 630 | ||||
| -rw-r--r-- | src/lib.rs | 26 | ||||
| -rw-r--r-- | src/main.rs | 39 |
5 files changed, 758 insertions, 1 deletions
diff --git a/src/configuration.rs b/src/configuration.rs new file mode 100644 index 0000000..f74d96f --- /dev/null +++ b/src/configuration.rs @@ -0,0 +1,44 @@ +use lexopt::{Parser, prelude::*}; + +const DEFAULT_ADDRESS: &str = "127.0.0.1:5500"; + +#[derive(Clone)] +pub struct Configuration { + pub address: String, +} + +impl Default for Configuration { + fn default() -> Self { + Self::new() + } +} + +impl Configuration { + #[must_use] + pub fn new() -> Self { + let mut address = DEFAULT_ADDRESS.to_string(); + + let mut parser = Parser::from_env(); + + while let Ok(Some(argument)) = parser.next() { + match argument { + Short('l') | Long("listen-address") => { + if let Ok(value) = parser.value().and_then(|v| v.parse()) { + address = value; + } else { + eprintln!("Warning: Invalid listen address ignored."); + } + } + Long("help") => { + println!("Usage: linea-caliente [-l|--listen-address=LISTEN_ADDRESS]"); + std::process::exit(0); + } + _ => { + eprintln!("Warning: Unknown argument ignored"); + } + } + } + + Configuration { address } + } +} diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..c07b418 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,20 @@ +use std::io; +use thiserror::Error; + +#[derive(Error, Debug)] +pub enum HotlineError { + #[error("client disconnected")] + Disconnected(#[from] io::Error), + #[error("handshake failed")] + HandshakeFailed, + #[error("invalid protocol ID")] + InvalidProtocolId, + #[error("unsupported version: {0}")] + UnsupportedVersion(u16), + #[error("incomplete handshake data")] + IncompleteHandshake, + #[error("server rejected handshake")] + ServerRejected, + #[error("invalid server response")] + InvalidServerResponse, +} diff --git a/src/handshake.rs b/src/handshake.rs new file mode 100644 index 0000000..9927181 --- /dev/null +++ b/src/handshake.rs @@ -0,0 +1,630 @@ +use crate::error::HotlineError; +use std::io::{Read, Write}; + +const PROTOCOL_ID: &[u8; 4] = b"TRTP"; +const HANDSHAKE_REQUEST_SIZE: usize = 12; +const HANDSHAKE_RESPONSE_SIZE: usize = 8; + +#[derive(Debug, Clone, PartialEq)] +pub struct HandshakeRequest { + pub protocol_id: [u8; 4], + pub sub_protocol_id: u32, + pub version: u16, + pub sub_version: u16, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct HandshakeResponse { + pub protocol_id: [u8; 4], + pub error_code: u32, +} + +impl HandshakeRequest { + #[must_use] + pub fn new(sub_protocol_id: u32, sub_version: u16) -> Self { + HandshakeRequest { + protocol_id: *PROTOCOL_ID, + sub_protocol_id, + version: 1, + sub_version, + } + } + + /// # Errors + /// + /// Returns error if unable to read handshake data from stream + pub fn parse(mut reader: impl Read) -> Result<Self, HotlineError> { + let mut buffer = [0u8; HANDSHAKE_REQUEST_SIZE]; + reader.read_exact(&mut buffer).map_err(|e| match e.kind() { + std::io::ErrorKind::UnexpectedEof => HotlineError::IncompleteHandshake, + _ => HotlineError::Disconnected(e), + })?; + + let protocol_id = [buffer[0], buffer[1], buffer[2], buffer[3]]; + let sub_protocol_id = u32::from_be_bytes([buffer[4], buffer[5], buffer[6], buffer[7]]); + let version = u16::from_be_bytes([buffer[8], buffer[9]]); + let sub_version = u16::from_be_bytes([buffer[10], buffer[11]]); + + Ok(HandshakeRequest { + protocol_id, + sub_protocol_id, + version, + sub_version, + }) + } + + #[must_use] + pub fn is_valid(&self) -> bool { + self.protocol_id == *PROTOCOL_ID && self.version == 1 + } + + /// Validates the handshake request and returns specific error information. + /// + /// # Errors + /// + /// Returns specific validation errors for debugging + pub fn validate(&self) -> Result<(), HotlineError> { + if self.protocol_id != *PROTOCOL_ID { + return Err(HotlineError::InvalidProtocolId); + } + + if self.version != 1 { + return Err(HotlineError::UnsupportedVersion(self.version)); + } + + Ok(()) + } + + /// # Errors + /// + /// Returns error if unable to write handshake data to stream + pub fn write_to(self, mut writer: impl Write) -> Result<(), HotlineError> { + let mut buffer = [0u8; HANDSHAKE_REQUEST_SIZE]; + + buffer[0..4].copy_from_slice(&self.protocol_id); + buffer[4..8].copy_from_slice(&self.sub_protocol_id.to_be_bytes()); + buffer[8..10].copy_from_slice(&self.version.to_be_bytes()); + buffer[10..12].copy_from_slice(&self.sub_version.to_be_bytes()); + + writer.write_all(&buffer)?; + writer.flush()?; + + Ok(()) + } +} + +impl HandshakeResponse { + #[must_use] + pub fn new(error_code: u32) -> Self { + HandshakeResponse { + protocol_id: *PROTOCOL_ID, + error_code, + } + } + + #[must_use] + pub fn success() -> Self { + Self::new(0) + } + + #[must_use] + pub fn error() -> Self { + Self::new(1) + } + + /// # Errors + /// + /// Returns error if unable to read handshake response from stream + pub fn parse(mut reader: impl Read) -> Result<Self, HotlineError> { + let mut buffer = [0u8; HANDSHAKE_RESPONSE_SIZE]; + reader.read_exact(&mut buffer).map_err(|e| match e.kind() { + std::io::ErrorKind::UnexpectedEof => HotlineError::IncompleteHandshake, + _ => HotlineError::Disconnected(e), + })?; + + let protocol_id = [buffer[0], buffer[1], buffer[2], buffer[3]]; + let error_code = u32::from_be_bytes([buffer[4], buffer[5], buffer[6], buffer[7]]); + + Ok(HandshakeResponse { + protocol_id, + error_code, + }) + } + + #[must_use] + pub fn is_success(&self) -> bool { + self.protocol_id == *PROTOCOL_ID && self.error_code == 0 + } + + /// # Errors + /// + /// Returns error if unable to write handshake response to stream + pub fn write_to(self, mut writer: impl Write) -> Result<(), HotlineError> { + let mut buffer = [0u8; HANDSHAKE_RESPONSE_SIZE]; + + buffer[0..4].copy_from_slice(&self.protocol_id); + buffer[4..8].copy_from_slice(&self.error_code.to_be_bytes()); + + writer.write_all(&buffer)?; + writer.flush()?; + + Ok(()) + } +} + +/// # Errors +/// +/// Returns error if handshake fails or stream I/O fails +pub fn accept_handshake(mut stream: impl Read + Write) -> Result<HandshakeRequest, HotlineError> { + let request = HandshakeRequest::parse(&mut stream)?; + + let validation_result = request.validate(); + + let response = match validation_result { + Ok(()) => HandshakeResponse::success(), + Err(_) => HandshakeResponse::error(), + }; + + response.write_to(&mut stream)?; + + validation_result?; + + Ok(request) +} + +/// # Errors +/// +/// Returns error if handshake is rejected by server or stream I/O fails +pub fn initiate_handshake( + mut stream: impl Read + Write, + sub_protocol_id: u32, + sub_version: u16, +) -> Result<(), HotlineError> { + let request = HandshakeRequest::new(sub_protocol_id, sub_version); + request.write_to(&mut stream)?; + + let response = HandshakeResponse::parse(&mut stream)?; + + if response.protocol_id != *PROTOCOL_ID { + return Err(HotlineError::InvalidServerResponse); + } + + if response.error_code != 0 { + return Err(HotlineError::ServerRejected); + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::io::Cursor; + + #[test] + fn test_valid_handshake_request() { + let mut data = Vec::new(); + data.extend_from_slice(b"TRTP"); + data.extend_from_slice(&0x1234_5678_u32.to_be_bytes()); + data.extend_from_slice(&1u16.to_be_bytes()); + data.extend_from_slice(&0x9ABCu16.to_be_bytes()); + + let request = HandshakeRequest::parse(Cursor::new(data)).unwrap(); + + assert_eq!(request.protocol_id, *b"TRTP"); + assert_eq!(request.sub_protocol_id, 0x1234_5678); + assert_eq!(request.version, 1); + assert_eq!(request.sub_version, 0x9ABC); + assert!(request.is_valid()); + } + + #[test] + fn test_invalid_protocol_id() { + let mut data = Vec::new(); + data.extend_from_slice(b"XXXX"); + data.extend_from_slice(&0u32.to_be_bytes()); + data.extend_from_slice(&1u16.to_be_bytes()); + data.extend_from_slice(&0u16.to_be_bytes()); + + let request = HandshakeRequest::parse(Cursor::new(data)).unwrap(); + assert!(!request.is_valid()); + } + + #[test] + fn test_handshake_response_write() { + let response = HandshakeResponse::success(); + let mut buffer = Vec::new(); + + response.write_to(&mut buffer).unwrap(); + + assert_eq!(buffer.len(), 8); + assert_eq!(&buffer[0..4], b"TRTP"); + assert_eq!( + u32::from_be_bytes([buffer[4], buffer[5], buffer[6], buffer[7]]), + 0 + ); + } + + #[test] + fn test_full_handshake_success() { + // Use a wrapper that implements both Read and Write + struct TestStream { + read_data: Cursor<Vec<u8>>, + write_data: Vec<u8>, + } + + impl Read for TestStream { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + self.read_data.read(buf) + } + } + + impl Write for TestStream { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + self.write_data.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + let mut request_data = Vec::new(); + request_data.extend_from_slice(b"TRTP"); + request_data.extend_from_slice(&0u32.to_be_bytes()); + request_data.extend_from_slice(&1u16.to_be_bytes()); + request_data.extend_from_slice(&0u16.to_be_bytes()); + + let stream = Cursor::new(request_data); + + let mut test_stream = TestStream { + read_data: stream, + write_data: Vec::new(), + }; + + let result = accept_handshake(&mut test_stream); + assert!(result.is_ok()); + + let request = result.unwrap(); + assert_eq!(request.protocol_id, *b"TRTP"); + assert_eq!(request.version, 1); + + // Check response was written + assert_eq!(test_stream.write_data.len(), 8); + assert_eq!(&test_stream.write_data[0..4], b"TRTP"); + assert_eq!( + u32::from_be_bytes([ + test_stream.write_data[4], + test_stream.write_data[5], + test_stream.write_data[6], + test_stream.write_data[7] + ]), + 0 + ); + } + + #[test] + fn test_handshake_request_new() { + let request = HandshakeRequest::new(0x1234_5678, 0x9ABC); + + assert_eq!(request.protocol_id, *b"TRTP"); + assert_eq!(request.sub_protocol_id, 0x1234_5678); + assert_eq!(request.version, 1); + assert_eq!(request.sub_version, 0x9ABC); + assert!(request.is_valid()); + } + + #[test] + fn test_handshake_request_write() { + let request = HandshakeRequest::new(0x1234_5678, 0x9ABC); + let mut buffer = Vec::new(); + + request.write_to(&mut buffer).unwrap(); + + assert_eq!(buffer.len(), 12); + assert_eq!(&buffer[0..4], b"TRTP"); + assert_eq!( + u32::from_be_bytes([buffer[4], buffer[5], buffer[6], buffer[7]]), + 0x1234_5678 + ); + assert_eq!(u16::from_be_bytes([buffer[8], buffer[9]]), 1); + assert_eq!(u16::from_be_bytes([buffer[10], buffer[11]]), 0x9ABC); + } + + #[test] + fn test_handshake_response_parse() { + let mut data = Vec::new(); + data.extend_from_slice(b"TRTP"); + data.extend_from_slice(&0u32.to_be_bytes()); + + let response = HandshakeResponse::parse(Cursor::new(data)).unwrap(); + + assert_eq!(response.protocol_id, *b"TRTP"); + assert_eq!(response.error_code, 0); + assert!(response.is_success()); + } + + #[test] + fn test_handshake_response_parse_error() { + let mut data = Vec::new(); + data.extend_from_slice(b"TRTP"); + data.extend_from_slice(&1u32.to_be_bytes()); + + let response = HandshakeResponse::parse(Cursor::new(data)).unwrap(); + + assert_eq!(response.protocol_id, *b"TRTP"); + assert_eq!(response.error_code, 1); + assert!(!response.is_success()); + } + + #[test] + fn test_initiate_handshake_success() { + struct MockStream { + read_data: Cursor<Vec<u8>>, + write_data: Vec<u8>, + } + + impl Read for MockStream { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + self.read_data.read(buf) + } + } + + impl Write for MockStream { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + self.write_data.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + // Prepare server response (success) + let mut response_data = Vec::new(); + response_data.extend_from_slice(b"TRTP"); + response_data.extend_from_slice(&0u32.to_be_bytes()); + + let mut mock_stream = MockStream { + read_data: Cursor::new(response_data), + write_data: Vec::new(), + }; + + let result = initiate_handshake(&mut mock_stream, 0x1234_5678, 0x9ABC); + assert!(result.is_ok()); + + // Check that request was sent correctly + assert_eq!(mock_stream.write_data.len(), 12); + assert_eq!(&mock_stream.write_data[0..4], b"TRTP"); + assert_eq!( + u32::from_be_bytes([ + mock_stream.write_data[4], + mock_stream.write_data[5], + mock_stream.write_data[6], + mock_stream.write_data[7] + ]), + 0x1234_5678 + ); + assert_eq!( + u16::from_be_bytes([mock_stream.write_data[8], mock_stream.write_data[9]]), + 1 + ); + assert_eq!( + u16::from_be_bytes([mock_stream.write_data[10], mock_stream.write_data[11]]), + 0x9ABC + ); + } + + #[test] + fn test_initiate_handshake_server_rejection() { + struct MockStream { + read_data: Cursor<Vec<u8>>, + write_data: Vec<u8>, + } + + impl Read for MockStream { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + self.read_data.read(buf) + } + } + + impl Write for MockStream { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + self.write_data.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + // Prepare server response (error) + let mut response_data = Vec::new(); + response_data.extend_from_slice(b"TRTP"); + response_data.extend_from_slice(&1u32.to_be_bytes()); + + let mut mock_stream = MockStream { + read_data: Cursor::new(response_data), + write_data: Vec::new(), + }; + + let result = initiate_handshake(&mut mock_stream, 0x1234_5678, 0x9ABC); + assert!(result.is_err()); + + match result.unwrap_err() { + HotlineError::ServerRejected => {} + _ => panic!("Expected ServerRejected error"), + } + } + + #[test] + fn test_client_server_interaction() { + // Test individual components work together correctly + // First, simulate what a client sends + let request = HandshakeRequest::new(0x1234_5678, 0x1234); + let mut request_buffer = Vec::new(); + request.write_to(&mut request_buffer).unwrap(); + + // Then test server can parse it and respond + let parsed_request = HandshakeRequest::parse(Cursor::new(&request_buffer)).unwrap(); + assert_eq!(parsed_request.sub_protocol_id, 0x1234_5678); + assert_eq!(parsed_request.sub_version, 0x1234); + assert!(parsed_request.is_valid()); + + // Server creates response + let response = HandshakeResponse::success(); + let mut response_buffer = Vec::new(); + response.write_to(&mut response_buffer).unwrap(); + + // Client can parse server response + let parsed_response = HandshakeResponse::parse(Cursor::new(&response_buffer)).unwrap(); + assert!(parsed_response.is_success()); + } + + #[test] + fn test_invalid_protocol_id_specific_error() { + let mut data = Vec::new(); + data.extend_from_slice(b"XXXX"); // Invalid protocol + data.extend_from_slice(&0u32.to_be_bytes()); + data.extend_from_slice(&1u16.to_be_bytes()); + data.extend_from_slice(&0u16.to_be_bytes()); + + let request = HandshakeRequest::parse(Cursor::new(data)).unwrap(); + + match request.validate().unwrap_err() { + HotlineError::InvalidProtocolId => {} + _ => panic!("Expected InvalidProtocolId error"), + } + } + + #[test] + fn test_unsupported_version_specific_error() { + let mut data = Vec::new(); + data.extend_from_slice(b"TRTP"); + data.extend_from_slice(&0u32.to_be_bytes()); + data.extend_from_slice(&2u16.to_be_bytes()); // Version 2 instead of 1 + data.extend_from_slice(&0u16.to_be_bytes()); + + let request = HandshakeRequest::parse(Cursor::new(data)).unwrap(); + + match request.validate().unwrap_err() { + HotlineError::UnsupportedVersion(version) => assert_eq!(version, 2), + _ => panic!("Expected UnsupportedVersion error"), + } + } + + #[test] + fn test_incomplete_handshake_request() { + // Only 8 bytes instead of 12 + let incomplete_data = vec![0x54, 0x52, 0x54, 0x50, 0x00, 0x00, 0x00, 0x00]; + + let result = HandshakeRequest::parse(Cursor::new(incomplete_data)); + + match result.unwrap_err() { + HotlineError::IncompleteHandshake => {} + _ => panic!("Expected IncompleteHandshake error"), + } + } + + #[test] + fn test_incomplete_handshake_response() { + // Only 4 bytes instead of 8 + let incomplete_data = vec![0x54, 0x52, 0x54, 0x50]; + + let result = HandshakeResponse::parse(Cursor::new(incomplete_data)); + + match result.unwrap_err() { + HotlineError::IncompleteHandshake => {} + _ => panic!("Expected IncompleteHandshake error"), + } + } + + #[test] + fn test_invalid_server_response() { + struct MockStream { + read_data: Cursor<Vec<u8>>, + write_data: Vec<u8>, + } + + impl Read for MockStream { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + self.read_data.read(buf) + } + } + + impl Write for MockStream { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + self.write_data.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + // Invalid server response with wrong protocol ID + let mut response_data = Vec::new(); + response_data.extend_from_slice(b"XXXX"); // Wrong protocol + response_data.extend_from_slice(&0u32.to_be_bytes()); + + let mut mock_stream = MockStream { + read_data: Cursor::new(response_data), + write_data: Vec::new(), + }; + + let result = initiate_handshake(&mut mock_stream, 0x1234_5678, 0x9ABC); + + match result.unwrap_err() { + HotlineError::InvalidServerResponse => {} + _ => panic!("Expected InvalidServerResponse error"), + } + } + + #[test] + fn test_accept_handshake_with_specific_errors() { + struct TestStream { + read_data: Cursor<Vec<u8>>, + write_data: Vec<u8>, + } + + impl Read for TestStream { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + self.read_data.read(buf) + } + } + + impl Write for TestStream { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + self.write_data.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + + // Test with invalid protocol ID + let mut request_data = Vec::new(); + request_data.extend_from_slice(b"XXXX"); + request_data.extend_from_slice(&0u32.to_be_bytes()); + request_data.extend_from_slice(&1u16.to_be_bytes()); + request_data.extend_from_slice(&0u16.to_be_bytes()); + + let mut test_stream = TestStream { + read_data: Cursor::new(request_data), + write_data: Vec::new(), + }; + + let result = accept_handshake(&mut test_stream); + + match result.unwrap_err() { + HotlineError::InvalidProtocolId => {} + _ => panic!("Expected InvalidProtocolId error"), + } + } +} @@ -1,2 +1,26 @@ -pub mod field; +pub mod error; +pub mod handshake; +// pub mod field; // pub mod transactions; + +use crate::error::HotlineError; +use crate::handshake::accept_handshake; + +use std::net::TcpStream; + +/// Get the stream and act as a server. If you want to act as a client, you +/// will have to write your own handler. +/// +/// # Errors +/// +/// Returns error if handshake fails or stream I/O fails +pub fn handle_client(mut stream: TcpStream) -> Result<(), HotlineError> { + let handshake_request = accept_handshake(&mut stream)?; + + eprintln!( + "Handshake successful with sub_protocol_id: 0x{:08X}, sub_version: {}", + handshake_request.sub_protocol_id, handshake_request.sub_version + ); + + Ok(()) +} diff --git a/src/main.rs b/src/main.rs new file mode 100644 index 0000000..d676896 --- /dev/null +++ b/src/main.rs @@ -0,0 +1,39 @@ +// These modules are "Server" modules. They're not necessary for the +// library, and we expect the library clients handle it as makes sense +// in their own processes. +pub mod configuration; + +use configuration::Configuration; +use linea_caliente::handle_client; + +use std::io::Result; +use std::net::TcpListener; +use std::thread; + +/// Spawns a server and hands over connections to the hotline library. +fn main() -> Result<()> { + let configuration = Configuration::new(); + + let listener = TcpListener::bind(&configuration.address)?; + eprintln!( + "Linea Caliente listening on address {}.", + configuration.address + ); + + for stream in listener.incoming() { + match stream { + Ok(stream) => { + thread::spawn(move || { + if let Err(error) = handle_client(stream) { + eprintln!("Error handling client: {error}"); + } + }); + } + Err(error) => { + eprintln!("Connection failed: {error}"); + } + } + } + + Ok(()) +} |