Async handshake
This commit is contained in:
@@ -0,0 +1,77 @@
|
||||
use tokio::{
|
||||
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
|
||||
net::TcpStream,
|
||||
};
|
||||
|
||||
pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> {
|
||||
let (read_half, mut write_half) = stream.split();
|
||||
let mut reader = BufReader::new(read_half);
|
||||
|
||||
let mut request_line = String::new();
|
||||
reader.read_line(&mut request_line).await?;
|
||||
let request_line = request_line.trim_end();
|
||||
|
||||
if request_line.starts_with("HEAD") {
|
||||
write_half
|
||||
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n")
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if !request_line.starts_with("GET") {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"Invalid HTTP method",
|
||||
));
|
||||
}
|
||||
|
||||
use std::collections::HashMap;
|
||||
let mut headers = HashMap::new();
|
||||
let mut line = String::new();
|
||||
|
||||
loop {
|
||||
line.clear();
|
||||
reader.read_line(&mut line).await?;
|
||||
if line == "\r\n" {
|
||||
break;
|
||||
}
|
||||
if let Some((k, v)) = line.split_once(':') {
|
||||
headers.insert(k.trim().to_lowercase(), v.trim().to_string());
|
||||
}
|
||||
}
|
||||
|
||||
if headers
|
||||
.get("upgrade")
|
||||
.map(|v| !v.eq_ignore_ascii_case("websocket"))
|
||||
.unwrap_or(true)
|
||||
{
|
||||
write_half
|
||||
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK")
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let key = headers
|
||||
.get("sec-websocket-key")
|
||||
.ok_or_else(|| std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing key"))?;
|
||||
|
||||
use base64::Engine;
|
||||
use base64::engine::general_purpose::STANDARD as Base64;
|
||||
use sha1::{Digest, Sha1};
|
||||
|
||||
let mut hasher = Sha1::new();
|
||||
hasher.update(key.as_bytes());
|
||||
hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
|
||||
let accept = Base64.encode(hasher.finalize());
|
||||
|
||||
let response = format!(
|
||||
"HTTP/1.1 101 Switching Protocols\r\n\
|
||||
Upgrade: websocket\r\n\
|
||||
Connection: Upgrade\r\n\
|
||||
Sec-WebSocket-Accept: {}\r\n\r\n",
|
||||
accept
|
||||
);
|
||||
|
||||
write_half.write_all(response.as_bytes()).await?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
pub mod handshake;
|
||||
pub mod server;
|
||||
pub mod session;
|
||||
|
||||
|
||||
+1
-113
@@ -8,124 +8,12 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::SessionMessage;
|
||||
|
||||
pub mod handshake {
|
||||
use base64::Engine;
|
||||
use base64::engine::general_purpose::STANDARD as Base64;
|
||||
use sha1::{Digest, Sha1};
|
||||
use std::collections::HashMap;
|
||||
use std::io::{BufRead, BufReader, Write};
|
||||
use std::net::TcpStream;
|
||||
|
||||
const WS_GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
|
||||
|
||||
pub fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Result<()> {
|
||||
let mut reader = BufReader::new(stream.try_clone()?);
|
||||
let mut request_line = String::new();
|
||||
reader.read_line(&mut request_line)?;
|
||||
|
||||
// Trim CRLF to make sure comparisons are clean
|
||||
let request_line = request_line.trim_end();
|
||||
|
||||
// Allow HEAD (used by Render for health checks)
|
||||
if request_line.starts_with("HEAD") {
|
||||
let response = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n";
|
||||
stream.write_all(response.as_bytes())?;
|
||||
stream.flush()?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Only proceed if it’s a GET
|
||||
if !request_line.starts_with("GET") {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
format!("Invalid HTTP method: {request_line}"),
|
||||
));
|
||||
}
|
||||
|
||||
// Read headers
|
||||
let mut headers = HashMap::new();
|
||||
let mut line = String::new();
|
||||
loop {
|
||||
line.clear();
|
||||
let bytes = reader.read_line(&mut line)?;
|
||||
if bytes == 0 || line == "\r\n" {
|
||||
break;
|
||||
}
|
||||
if let Some((k, v)) = line.split_once(':') {
|
||||
headers.insert(k.trim().to_lowercase(), v.trim().to_string());
|
||||
}
|
||||
}
|
||||
|
||||
// Check if it's actually a WebSocket upgrade request
|
||||
let is_websocket_upgrade = headers
|
||||
.get("upgrade")
|
||||
.map(|v| v.eq_ignore_ascii_case("websocket"))
|
||||
.unwrap_or(false);
|
||||
|
||||
if !is_websocket_upgrade {
|
||||
// Not a WebSocket request — probably a normal HTTP GET (e.g. health check)
|
||||
let response =
|
||||
"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\n\r\nOK";
|
||||
stream.write_all(response.as_bytes())?;
|
||||
stream.flush()?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Validate "Connection: Upgrade"
|
||||
if !headers
|
||||
.get("connection")
|
||||
.map(|v| v.to_lowercase().contains("upgrade"))
|
||||
.unwrap_or(false)
|
||||
{
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
"Missing or invalid Connection header",
|
||||
));
|
||||
}
|
||||
|
||||
// Validate WebSocket key
|
||||
let key = headers.get("sec-websocket-key").ok_or_else(|| {
|
||||
std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key")
|
||||
})?;
|
||||
|
||||
// Validate version
|
||||
if let Some(ver) = headers.get("sec-websocket-version") {
|
||||
if ver.trim() != "13" {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::InvalidData,
|
||||
format!("Unsupported Sec-WebSocket-Version: {}", ver),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
// Compute accept key
|
||||
let mut hasher = Sha1::new();
|
||||
hasher.update(key.as_bytes());
|
||||
hasher.update(WS_GUID.as_bytes());
|
||||
let hash = hasher.finalize();
|
||||
let accept_key = Base64.encode(hash);
|
||||
|
||||
// Send response
|
||||
let response = format!(
|
||||
"HTTP/1.1 101 Switching Protocols\r\n\
|
||||
Upgrade: websocket\r\n\
|
||||
Connection: Upgrade\r\n\
|
||||
Sec-WebSocket-Accept: {}\r\n\r\n",
|
||||
accept_key
|
||||
);
|
||||
|
||||
stream.write_all(response.as_bytes())?;
|
||||
stream.flush()?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Session(TcpStream, u64);
|
||||
|
||||
impl Session {
|
||||
/// Create a client
|
||||
pub fn new(mut stream: TcpStream) -> crate::Result<Self> {
|
||||
handshake::handle_websocket_handshake(&mut stream)?;
|
||||
crate::handshake::handle_websocket_handshake(&mut stream)?;
|
||||
stream.set_read_timeout(Some(std::time::Duration::from_secs(10)))?;
|
||||
stream.set_write_timeout(Some(std::time::Duration::from_secs(10)))?;
|
||||
Ok(Session(stream, rand::random()))
|
||||
|
||||
Reference in New Issue
Block a user