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 { 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 { 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 { 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>, write_data: Vec, } impl Read for TestStream { fn read(&mut self, buf: &mut [u8]) -> std::io::Result { self.read_data.read(buf) } } impl Write for TestStream { fn write(&mut self, buf: &[u8]) -> std::io::Result { 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>, write_data: Vec, } impl Read for MockStream { fn read(&mut self, buf: &mut [u8]) -> std::io::Result { self.read_data.read(buf) } } impl Write for MockStream { fn write(&mut self, buf: &[u8]) -> std::io::Result { 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>, write_data: Vec, } impl Read for MockStream { fn read(&mut self, buf: &mut [u8]) -> std::io::Result { self.read_data.read(buf) } } impl Write for MockStream { fn write(&mut self, buf: &[u8]) -> std::io::Result { 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>, write_data: Vec, } impl Read for MockStream { fn read(&mut self, buf: &mut [u8]) -> std::io::Result { self.read_data.read(buf) } } impl Write for MockStream { fn write(&mut self, buf: &[u8]) -> std::io::Result { 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>, write_data: Vec, } impl Read for TestStream { fn read(&mut self, buf: &mut [u8]) -> std::io::Result { self.read_data.read(buf) } } impl Write for TestStream { fn write(&mut self, buf: &[u8]) -> std::io::Result { 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"), } } }