diff --git a/examples/client.rs b/examples/client.rs index f743c59..84b4abb 100644 --- a/examples/client.rs +++ b/examples/client.rs @@ -10,14 +10,13 @@ async fn main() -> session_rs::Result<()> { let read_session = Arc::clone(&session); tokio::spawn(async move { loop { - println!("{:?}", read_session.read().await); - // match read_session.read().await { - // Ok(Frame::Text(text)) => { - // println!("Server says: {}", text); - // } - // Ok(_) => {} - // Err(_) => break, - // } + match read_session.read().await { + Ok(Frame::Text(text)) => { + println!("Server says: {}", text); + } + Ok(_) => {} + Err(_) => break, + } } }); diff --git a/src/ws/handshake.rs b/src/ws/handshake.rs index 4f92ec6..b9bfcb0 100644 --- a/src/ws/handshake.rs +++ b/src/ws/handshake.rs @@ -91,7 +91,7 @@ impl WebSocket { id: rand::random(), reader: Arc::new(Mutex::new(read)), writer: Arc::new(Mutex::new(write)), - mask_payload: false, + is_server: false, }) } @@ -165,7 +165,7 @@ impl WebSocket { id: rand::random(), reader: Arc::new(Mutex::new(read)), writer: Arc::new(Mutex::new(write)), - mask_payload: true, + is_server: true, }) } } diff --git a/src/ws/mod.rs b/src/ws/mod.rs index a72bf41..6f6227a 100644 --- a/src/ws/mod.rs +++ b/src/ws/mod.rs @@ -24,7 +24,7 @@ pub struct WebSocket { pub(crate) reader: Arc>, pub(crate) writer: Arc>, pub(crate) id: u64, - pub(crate) mask_payload: bool, + pub(crate) is_server: bool, } impl Clone for WebSocket { @@ -32,7 +32,7 @@ impl Clone for WebSocket { WebSocket { reader: self.reader.clone(), writer: self.writer.clone(), - mask_payload: self.mask_payload.clone(), + is_server: self.is_server.clone(), id: self.id, } } @@ -57,7 +57,7 @@ impl WebSocket { let mut writer = self.writer.lock().await; let mut header = Vec::with_capacity(10); - let mask_bit = if self.mask_payload { 0x80 } else { 0x00 }; + let mask_bit = if self.is_server { 0x80 } else { 0x00 }; header.push(0x80 | opcode); // FIN + opcode let len = payload.len(); @@ -71,7 +71,7 @@ impl WebSocket { header.extend_from_slice(&(len as u64).to_be_bytes()); } - if self.mask_payload { + if self.is_server { // Generate 4-byte mask key let mask_key: [u8; 4] = rand::random(); header.extend_from_slice(&mask_key); @@ -155,26 +155,33 @@ impl WebSocket { payload_len = u64::from_be_bytes(buf); } - // --- 3. Read mask key --- - if !masked && !self.mask_payload { - // Per spec, client-to-server frames MUST be masked - self.close().await.ok(); - return Err(Error::InvalidFrame( - "Received unmasked frame from client".into(), - )); - } - - let mut mask = [0u8; 4]; - reader.read_exact(&mut mask).await?; - - // --- 4. Read payload --- - let mut payload = vec![0u8; payload_len as usize]; - if payload_len > 0 { - reader.read_exact(&mut payload).await?; - for i in 0..payload.len() { - payload[i] ^= mask[i % 4]; + let payload = if masked { + // --- 3. Read mask key --- + let mut mask = [0u8; 4]; + reader.read_exact(&mut mask).await?; + let mut payload = vec![0u8; payload_len as usize]; + if payload_len > 0 { + reader.read_exact(&mut payload).await?; + for i in 0..payload.len() { + payload[i] ^= mask[i % 4]; + } } - } + payload + } else { + // Per spec, client-to-server frames MUST be masked + if !self.is_server { + self.close().await.ok(); + return Err(Error::InvalidFrame( + "Received unmasked frame from client".into(), + )); + } + + let mut payload = vec![0u8; payload_len as usize]; + if payload_len > 0 { + reader.read_exact(&mut payload).await?; + } + payload + }; // --- 6. Return opcode + payload --- Ok((fin, opcode, payload))