Refactored masking
This commit is contained in:
+7
-8
@@ -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
@@ -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
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user