Merge pull request #3 from selimaj-dev/fix-websocket-handling
fix: stabilize websocket session handling and eliminate deadlocks
This commit is contained in:
+2
-2
@@ -1,4 +1,4 @@
|
|||||||
use std::pin::Pin;
|
use std::{pin::Pin, sync::Arc};
|
||||||
|
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
@@ -9,7 +9,7 @@ pub mod ws;
|
|||||||
pub type Result<T> = std::result::Result<T, Error>;
|
pub type Result<T> = std::result::Result<T, Error>;
|
||||||
pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> =
|
pub type BoxFuture<'a, T = Option<(bool, serde_json::Value)>> =
|
||||||
Pin<Box<dyn Future<Output = T> + Send + 'a>>;
|
Pin<Box<dyn Future<Output = T> + Send + 'a>>;
|
||||||
pub type MethodHandler = Box<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>;
|
pub type MethodHandler = Arc<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>;
|
||||||
|
|
||||||
pub trait Method {
|
pub trait Method {
|
||||||
const NAME: &'static str;
|
const NAME: &'static str;
|
||||||
|
|||||||
+29
-8
@@ -1,6 +1,6 @@
|
|||||||
use std::{net::SocketAddr, sync::Arc};
|
use std::{net::SocketAddr, sync::Arc};
|
||||||
|
|
||||||
use tokio::net::TcpListener;
|
use tokio::{net::TcpListener, time::timeout};
|
||||||
|
|
||||||
use crate::{session::Session, ws::WebSocket};
|
use crate::{session::Session, ws::WebSocket};
|
||||||
|
|
||||||
@@ -23,19 +23,40 @@ impl SessionServer {
|
|||||||
Ok((Session::from_ws(ws), addr))
|
Ok((Session::from_ws(ws), addr))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn session_loop<Fut: Future<Output = crate::Result<()>> + Send + 'static>(
|
pub async fn session_loop<F, Fut>(&self, on_conn: F) -> crate::Result<()>
|
||||||
&self,
|
where
|
||||||
on_conn: impl Fn(Session, SocketAddr) -> Fut + 'static,
|
F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static,
|
||||||
) -> crate::Result<()> {
|
Fut: Future<Output = crate::Result<()>> + Send + 'static,
|
||||||
|
{
|
||||||
let conn_handler = Arc::new(on_conn);
|
let conn_handler = Arc::new(on_conn);
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
let (session, addr) = self.accept().await?;
|
let (stream, addr) = self.listener.accept().await?;
|
||||||
let conn_handler = conn_handler.clone();
|
let conn_handler = conn_handler.clone();
|
||||||
|
|
||||||
session.start_receiver();
|
tokio::spawn(async move {
|
||||||
|
match timeout(
|
||||||
|
tokio::time::Duration::from_secs(5),
|
||||||
|
WebSocket::handshake(stream),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
{
|
||||||
|
Ok(Ok(ws)) => {
|
||||||
|
let session = Session::from_ws(ws);
|
||||||
|
session.start_receiver();
|
||||||
|
|
||||||
tokio::spawn(conn_handler(session, addr));
|
if let Err(e) = conn_handler(session, addr).await {
|
||||||
|
eprintln!("Connection error: {:?}", e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Ok(Err(e)) => {
|
||||||
|
eprintln!("Handshake failed from {}: {:?}", addr, e);
|
||||||
|
}
|
||||||
|
Err(_) => {
|
||||||
|
eprintln!("Handshake failed from {}: Handshake Timeout", addr);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+7
-2
@@ -87,7 +87,12 @@ impl Session {
|
|||||||
|
|
||||||
match msg {
|
match msg {
|
||||||
Message::Request { id, method, data } => {
|
Message::Request { id, method, data } => {
|
||||||
if let Some(m) = s.methods.lock().await.get(&method) {
|
let handler = {
|
||||||
|
let methods = s.methods.lock().await;
|
||||||
|
methods.get(&method).cloned()
|
||||||
|
};
|
||||||
|
|
||||||
|
if let Some(m) = handler {
|
||||||
if let Some((err, res)) = (m)(id, data).await {
|
if let Some((err, res)) = (m)(id, data).await {
|
||||||
if err {
|
if err {
|
||||||
s.respond_error(id, res)
|
s.respond_error(id, res)
|
||||||
@@ -157,7 +162,7 @@ impl Session {
|
|||||||
|
|
||||||
self.methods.lock().await.insert(
|
self.methods.lock().await.insert(
|
||||||
M::NAME.to_string(),
|
M::NAME.to_string(),
|
||||||
Box::new(move |id, value| {
|
Arc::new(move |id, value| {
|
||||||
let handler = Arc::clone(&handler);
|
let handler = Arc::clone(&handler);
|
||||||
|
|
||||||
Box::pin(async move {
|
Box::pin(async move {
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ pub enum Error {
|
|||||||
HandshakeFailed(String),
|
HandshakeFailed(String),
|
||||||
Utf8(FromUtf8Error),
|
Utf8(FromUtf8Error),
|
||||||
ConnectionClosed,
|
ConnectionClosed,
|
||||||
|
Elapsed,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl From<std::io::Error> for Error {
|
impl From<std::io::Error> for Error {
|
||||||
@@ -22,3 +23,9 @@ impl From<FromUtf8Error> for Error {
|
|||||||
Self::Utf8(value)
|
Self::Utf8(value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl From<tokio::time::error::Elapsed> for Error {
|
||||||
|
fn from(_: tokio::time::error::Elapsed) -> Self {
|
||||||
|
Self::Elapsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+81
-28
@@ -1,10 +1,12 @@
|
|||||||
use base64::Engine;
|
use base64::Engine;
|
||||||
|
use base64::engine::general_purpose::STANDARD as Base64;
|
||||||
use sha1::{Digest, Sha1};
|
use sha1::{Digest, Sha1};
|
||||||
use std::sync::Arc;
|
use std::{collections::HashMap, sync::Arc};
|
||||||
use tokio::{
|
use tokio::{
|
||||||
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
|
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
|
||||||
net::TcpStream,
|
net::TcpStream,
|
||||||
sync::Mutex,
|
sync::Mutex,
|
||||||
|
time::{Duration, timeout},
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::WebSocket;
|
use super::WebSocket;
|
||||||
@@ -13,72 +15,119 @@ pub async fn handle_websocket_handshake(stream: &mut TcpStream) -> std::io::Resu
|
|||||||
let (read_half, mut write_half) = stream.split();
|
let (read_half, mut write_half) = stream.split();
|
||||||
let mut reader = BufReader::new(read_half);
|
let mut reader = BufReader::new(read_half);
|
||||||
|
|
||||||
|
// ---- 1. Read request line with timeout ----
|
||||||
let mut request_line = String::new();
|
let mut request_line = String::new();
|
||||||
reader.read_line(&mut request_line).await?;
|
timeout(Duration::from_secs(5), reader.read_line(&mut request_line)).await??;
|
||||||
|
|
||||||
let request_line = request_line.trim_end();
|
let request_line = request_line.trim_end();
|
||||||
|
|
||||||
if request_line.starts_with("HEAD") {
|
if !request_line.starts_with("GET") {
|
||||||
write_half
|
write_half
|
||||||
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n")
|
.write_all(
|
||||||
|
b"HTTP/1.1 405 Method Not Allowed\r\n\
|
||||||
|
Content-Length: 0\r\n\
|
||||||
|
Connection: close\r\n\r\n",
|
||||||
|
)
|
||||||
.await?;
|
.await?;
|
||||||
|
write_half.shutdown().await?;
|
||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
|
||||||
if !request_line.starts_with("GET") {
|
// ---- 2. Read headers with timeout ----
|
||||||
return Err(std::io::Error::new(
|
|
||||||
std::io::ErrorKind::InvalidData,
|
|
||||||
"Invalid HTTP method",
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
let mut headers = HashMap::new();
|
let mut headers = HashMap::new();
|
||||||
let mut line = String::new();
|
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
line.clear();
|
let mut line = String::new();
|
||||||
reader.read_line(&mut line).await?;
|
timeout(Duration::from_secs(5), reader.read_line(&mut line)).await??;
|
||||||
|
|
||||||
if line == "\r\n" {
|
if line == "\r\n" {
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some((k, v)) = line.split_once(':') {
|
if let Some((k, v)) = line.split_once(':') {
|
||||||
headers.insert(k.trim().to_lowercase(), v.trim().to_string());
|
headers.insert(k.trim().to_lowercase(), v.trim().to_string());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if headers
|
// ---- 3. Check if this is a WebSocket upgrade ----
|
||||||
|
let is_upgrade = headers
|
||||||
.get("upgrade")
|
.get("upgrade")
|
||||||
.map(|v| !v.eq_ignore_ascii_case("websocket"))
|
.map(|v| v.eq_ignore_ascii_case("websocket"))
|
||||||
.unwrap_or(true)
|
.unwrap_or(false);
|
||||||
{
|
|
||||||
|
let has_connection_upgrade = headers
|
||||||
|
.get("connection")
|
||||||
|
.map(|v| v.to_lowercase().contains("upgrade"))
|
||||||
|
.unwrap_or(false);
|
||||||
|
|
||||||
|
if !is_upgrade || !has_connection_upgrade {
|
||||||
|
// Normal HTTP response (important for browsers)
|
||||||
|
let body = b"OK";
|
||||||
|
|
||||||
write_half
|
write_half
|
||||||
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nOK")
|
.write_all(
|
||||||
|
format!(
|
||||||
|
"HTTP/1.1 200 OK\r\n\
|
||||||
|
Content-Type: text/plain\r\n\
|
||||||
|
Content-Length: {}\r\n\
|
||||||
|
Connection: close\r\n\
|
||||||
|
\r\n",
|
||||||
|
body.len()
|
||||||
|
)
|
||||||
|
.as_bytes(),
|
||||||
|
)
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
|
write_half.write_all(body).await?;
|
||||||
|
write_half.flush().await?;
|
||||||
|
write_half.shutdown().await?;
|
||||||
|
|
||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
|
|
||||||
let key = headers
|
// ---- 4. Validate required headers ----
|
||||||
.get("sec-websocket-key")
|
let key = headers.get("sec-websocket-key").ok_or_else(|| {
|
||||||
.ok_or_else(|| std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing key"))?;
|
std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key")
|
||||||
|
})?;
|
||||||
|
|
||||||
use base64::Engine;
|
let version_ok = headers
|
||||||
use base64::engine::general_purpose::STANDARD as Base64;
|
.get("sec-websocket-version")
|
||||||
use sha1::{Digest, Sha1};
|
.map(|v| v == "13")
|
||||||
|
.unwrap_or(false);
|
||||||
|
|
||||||
|
if !version_ok {
|
||||||
|
write_half
|
||||||
|
.write_all(
|
||||||
|
b"HTTP/1.1 426 Upgrade Required\r\n\
|
||||||
|
Sec-WebSocket-Version: 13\r\n\
|
||||||
|
Content-Length: 0\r\n\
|
||||||
|
Connection: close\r\n\r\n",
|
||||||
|
)
|
||||||
|
.await?;
|
||||||
|
write_half.shutdown().await?;
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- 5. Generate Sec-WebSocket-Accept ----
|
||||||
let mut hasher = Sha1::new();
|
let mut hasher = Sha1::new();
|
||||||
hasher.update(key.as_bytes());
|
hasher.update(key.as_bytes());
|
||||||
hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
|
hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
|
||||||
|
|
||||||
let accept = Base64.encode(hasher.finalize());
|
let accept = Base64.encode(hasher.finalize());
|
||||||
|
|
||||||
|
// ---- 6. Send upgrade response ----
|
||||||
let response = format!(
|
let response = format!(
|
||||||
"HTTP/1.1 101 Switching Protocols\r\n\
|
"HTTP/1.1 101 Switching Protocols\r\n\
|
||||||
Upgrade: websocket\r\n\
|
Upgrade: websocket\r\n\
|
||||||
Connection: Upgrade\r\n\
|
Connection: Upgrade\r\n\
|
||||||
Sec-WebSocket-Accept: {}\r\n\r\n",
|
Sec-WebSocket-Accept: {}\r\n\
|
||||||
|
\r\n",
|
||||||
accept
|
accept
|
||||||
);
|
);
|
||||||
|
|
||||||
write_half.write_all(response.as_bytes()).await?;
|
write_half.write_all(response.as_bytes()).await?;
|
||||||
|
write_half.flush().await?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,7 +171,11 @@ impl WebSocket {
|
|||||||
// 4. Read HTTP response
|
// 4. Read HTTP response
|
||||||
let mut reader = BufReader::new(&mut stream);
|
let mut reader = BufReader::new(&mut stream);
|
||||||
let mut status_line = String::new();
|
let mut status_line = String::new();
|
||||||
reader.read_line(&mut status_line).await?;
|
timeout(
|
||||||
|
tokio::time::Duration::from_secs(5),
|
||||||
|
reader.read_line(&mut status_line),
|
||||||
|
)
|
||||||
|
.await??;
|
||||||
if !status_line.starts_with("HTTP/1.1 101") {
|
if !status_line.starts_with("HTTP/1.1 101") {
|
||||||
return Err(super::Error::HandshakeFailed(format!(
|
return Err(super::Error::HandshakeFailed(format!(
|
||||||
"Expected 101 Switching Protocols, got: {}",
|
"Expected 101 Switching Protocols, got: {}",
|
||||||
|
|||||||
Reference in New Issue
Block a user