Refactored masking

This commit is contained in:
2026-02-19 00:43:56 +01:00
parent 8d67252800
commit 60e5a54261
3 changed files with 39 additions and 33 deletions
+7 -8
View File
@@ -10,14 +10,13 @@ async fn main() -> session_rs::Result<()> {
let read_session = Arc::clone(&session); let read_session = Arc::clone(&session);
tokio::spawn(async move { tokio::spawn(async move {
loop { loop {
println!("{:?}", read_session.read().await); match read_session.read().await {
// match read_session.read().await { Ok(Frame::Text(text)) => {
// Ok(Frame::Text(text)) => { println!("Server says: {}", text);
// println!("Server says: {}", text); }
// } Ok(_) => {}
// Ok(_) => {} Err(_) => break,
// Err(_) => break, }
// }
} }
}); });
+2 -2
View File
@@ -91,7 +91,7 @@ impl WebSocket {
id: rand::random(), id: rand::random(),
reader: Arc::new(Mutex::new(read)), reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)), writer: Arc::new(Mutex::new(write)),
mask_payload: false, is_server: false,
}) })
} }
@@ -165,7 +165,7 @@ impl WebSocket {
id: rand::random(), id: rand::random(),
reader: Arc::new(Mutex::new(read)), reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)), writer: Arc::new(Mutex::new(write)),
mask_payload: true, is_server: true,
}) })
} }
} }
+21 -14
View File
@@ -24,7 +24,7 @@ pub struct WebSocket {
pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>, pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>,
pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>, pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
pub(crate) id: u64, pub(crate) id: u64,
pub(crate) mask_payload: bool, pub(crate) is_server: bool,
} }
impl Clone for WebSocket { impl Clone for WebSocket {
@@ -32,7 +32,7 @@ impl Clone for WebSocket {
WebSocket { WebSocket {
reader: self.reader.clone(), reader: self.reader.clone(),
writer: self.writer.clone(), writer: self.writer.clone(),
mask_payload: self.mask_payload.clone(), is_server: self.is_server.clone(),
id: self.id, id: self.id,
} }
} }
@@ -57,7 +57,7 @@ impl WebSocket {
let mut writer = self.writer.lock().await; let mut writer = self.writer.lock().await;
let mut header = Vec::with_capacity(10); 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 header.push(0x80 | opcode); // FIN + opcode
let len = payload.len(); let len = payload.len();
@@ -71,7 +71,7 @@ impl WebSocket {
header.extend_from_slice(&(len as u64).to_be_bytes()); header.extend_from_slice(&(len as u64).to_be_bytes());
} }
if self.mask_payload { if self.is_server {
// Generate 4-byte mask key // Generate 4-byte mask key
let mask_key: [u8; 4] = rand::random(); let mask_key: [u8; 4] = rand::random();
header.extend_from_slice(&mask_key); header.extend_from_slice(&mask_key);
@@ -155,19 +155,10 @@ impl WebSocket {
payload_len = u64::from_be_bytes(buf); payload_len = u64::from_be_bytes(buf);
} }
let payload = if masked {
// --- 3. Read mask key --- // --- 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]; let mut mask = [0u8; 4];
reader.read_exact(&mut mask).await?; reader.read_exact(&mut mask).await?;
// --- 4. Read payload ---
let mut payload = vec![0u8; payload_len as usize]; let mut payload = vec![0u8; payload_len as usize];
if payload_len > 0 { if payload_len > 0 {
reader.read_exact(&mut payload).await?; reader.read_exact(&mut payload).await?;
@@ -175,6 +166,22 @@ impl WebSocket {
payload[i] ^= mask[i % 4]; 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 --- // --- 6. Return opcode + payload ---
Ok((fin, opcode, payload)) Ok((fin, opcode, payload))