Make the protocol transport-agnostic and use tokio-tungstenite
Replace the hand-rolled WebSocket implementation with a transport boundary (Frame over any Sink/Stream) plus adapters for tokio-tungstenite (server/client features) and axum (axum feature). The JSON wire format is unchanged, so existing peers keep working. Fixes: - unbounded frame lengths were allocated up front; the server now enforces message/frame size limits (1 MiB default) - responses arriving before `request` subscribed were lost - an accept() error ended `session_loop` - requests sent right after connecting could arrive before handlers were registered; the receiver now starts after `on_conn` returns - panics in the receive loop skipped `on_close` and leaked sessions; `on_close` now runs exactly once and handler panics fail only their request - unknown methods and invalid data got no reply; they now get an error response - slow handlers blocked pongs and responses Adds `on_notification`, `request_timeout`, `closed`, `is_closed`, `id`, `ServerConfig`, `from_transport`, `from_tungstenite` and `from_axum`, integration tests, and an axum example. Bumps to 0.2.0 since `Session::connect` now takes a URL and the `ws` module is removed. Co-Authored-By: Claude Opus 5.5 <[email protected]>
This commit is contained in:
@@ -0,0 +1,452 @@
|
||||
#![cfg(all(feature = "server", feature = "client"))]
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use session_rs::server::SessionServer;
|
||||
use session_rs::{Error, Method, Session};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_tungstenite::tungstenite::Message as WsMessage;
|
||||
|
||||
macro_rules! method {
|
||||
($ty:ident, $name:literal, $req:ty, $res:ty) => {
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct $ty;
|
||||
|
||||
impl Method for $ty {
|
||||
const NAME: &'static str = $name;
|
||||
type Request = $req;
|
||||
type Response = $res;
|
||||
type Error = String;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
method!(Echo, "echo", String, String);
|
||||
method!(Fail, "fail", String, String);
|
||||
method!(Panic, "panic", (), ());
|
||||
method!(Numbers, "numbers", Vec<u32>, u32);
|
||||
method!(AskBack, "ask_back", String, String);
|
||||
method!(Notice, "notice", String, ());
|
||||
method!(Silent, "silent", (), ());
|
||||
|
||||
const WAIT: Duration = Duration::from_secs(5);
|
||||
|
||||
async fn register_handlers(session: &Session) {
|
||||
session.on_request::<Echo, _>(async |_, s| Ok(s)).await;
|
||||
session.on_request::<Fail, _>(async |_, s| Err(format!("failed: {s}"))).await;
|
||||
session
|
||||
.on_request::<Panic, _>(async |_, ()| -> Result<(), String> { panic!("boom") })
|
||||
.await;
|
||||
session
|
||||
.on_request::<Numbers, _>(async |_, v| Ok(v.iter().sum()))
|
||||
.await;
|
||||
session
|
||||
.on_request::<AskBack, _>({
|
||||
let session = session.clone();
|
||||
move |_, s| {
|
||||
let session = session.clone();
|
||||
async move {
|
||||
// A handler requesting from its own peer must not deadlock.
|
||||
let reply = session.request::<Echo>(format!("back:{s}")).await;
|
||||
reply.map_err(|e| e.to_string())?
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
session
|
||||
.on_request::<Silent, _>(async |_, ()| {
|
||||
tokio::time::sleep(Duration::from_secs(60)).await;
|
||||
Ok(())
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
async fn start_server() -> SocketAddr {
|
||||
start_server_with(|_| {}).await
|
||||
}
|
||||
|
||||
/// Starts a server; `on_session` sees every accepted session after its
|
||||
/// handlers are registered.
|
||||
async fn start_server_with(on_session: impl Fn(Session) + Send + Sync + 'static) -> SocketAddr {
|
||||
let server = SessionServer::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = server.local_addr().unwrap();
|
||||
let on_session = Arc::new(on_session);
|
||||
|
||||
tokio::spawn(async move {
|
||||
server
|
||||
.session_loop(move |session, _| {
|
||||
let on_session = on_session.clone();
|
||||
async move {
|
||||
// Registering late must not lose a request sent right after connecting.
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
register_handlers(&session).await;
|
||||
on_session(session);
|
||||
Ok(())
|
||||
}
|
||||
})
|
||||
.await
|
||||
});
|
||||
|
||||
addr
|
||||
}
|
||||
|
||||
async fn connect(addr: SocketAddr) -> Session {
|
||||
let session = Session::connect(&format!("ws://{addr}")).await.unwrap();
|
||||
register_handlers(&session).await;
|
||||
session.start_receiver();
|
||||
session
|
||||
}
|
||||
|
||||
async fn raw_client(addr: SocketAddr) -> tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>> {
|
||||
tokio_tungstenite::connect_async(format!("ws://{addr}")).await.unwrap().0
|
||||
}
|
||||
|
||||
async fn next_text(
|
||||
ws: &mut tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
|
||||
) -> serde_json::Value {
|
||||
loop {
|
||||
match tokio::time::timeout(WAIT, ws.next()).await.unwrap().unwrap().unwrap() {
|
||||
WsMessage::Text(t) => return serde_json::from_str(&t).unwrap(),
|
||||
_ => continue,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_response_and_error() {
|
||||
let client = connect(start_server().await).await;
|
||||
|
||||
assert_eq!(client.request::<Echo>("hi".into()).await.unwrap(), Ok("hi".into()));
|
||||
assert_eq!(client.request::<Numbers>(vec![1, 2, 3]).await.unwrap(), Ok(6));
|
||||
assert_eq!(
|
||||
client.request::<Fail>("x".into()).await.unwrap(),
|
||||
Err("failed: x".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_requests_are_matched_by_id() {
|
||||
let client = connect(start_server().await).await;
|
||||
|
||||
let replies = futures_util::future::join_all(
|
||||
(0..200).map(|i| {
|
||||
let client = client.clone();
|
||||
async move { client.request::<Echo>(i.to_string()).await.unwrap() }
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
for (i, reply) in replies.into_iter().enumerate() {
|
||||
assert_eq!(reply, Ok(i.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn wire_format_matches_protocol() {
|
||||
let mut ws = raw_client(start_server().await).await;
|
||||
|
||||
ws.send(WsMessage::Text(
|
||||
r#"{"type":"request","id":7,"method":"echo","data":"hi"}"#.into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
next_text(&mut ws).await,
|
||||
serde_json::json!({"type": "response", "id": 7, "result": "hi"})
|
||||
);
|
||||
|
||||
ws.send(WsMessage::Text(
|
||||
r#"{"type":"request","id":8,"method":"fail","data":"x"}"#.into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
next_text(&mut ws).await,
|
||||
serde_json::json!({"type": "errorresponse", "id": 8, "error": "failed: x"})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unknown_method_and_bad_data_get_error_responses() {
|
||||
let mut ws = raw_client(start_server().await).await;
|
||||
|
||||
ws.send(WsMessage::Text(
|
||||
r#"{"type":"request","id":1,"method":"nope","data":null}"#.into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
next_text(&mut ws).await,
|
||||
serde_json::json!({"type": "errorresponse", "id": 1, "error": "Unknown method: nope"})
|
||||
);
|
||||
|
||||
ws.send(WsMessage::Text(
|
||||
r#"{"type":"request","id":2,"method":"numbers","data":"not a list"}"#.into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let reply = next_text(&mut ws).await;
|
||||
assert_eq!(reply["type"], "errorresponse");
|
||||
assert_eq!(reply["id"], 2);
|
||||
assert!(reply["error"].as_str().unwrap().starts_with("Invalid request data"));
|
||||
|
||||
// Garbage text is ignored, not fatal.
|
||||
ws.send(WsMessage::Text("not json".into())).await.unwrap();
|
||||
ws.send(WsMessage::Text(
|
||||
r#"{"type":"request","id":3,"method":"echo","data":"still here"}"#.into(),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(next_text(&mut ws).await["result"], "still here");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn panicking_handler_fails_only_that_request() {
|
||||
let client = connect(start_server().await).await;
|
||||
|
||||
assert_eq!(
|
||||
client.request::<Panic>(()).await.unwrap(),
|
||||
Err("Handler panicked".into())
|
||||
);
|
||||
assert_eq!(client.request::<Echo>("ok".into()).await.unwrap(), Ok("ok".into()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn handler_can_request_from_its_peer() {
|
||||
let client = connect(start_server().await).await;
|
||||
|
||||
assert_eq!(
|
||||
client.request::<AskBack>("x".into()).await.unwrap(),
|
||||
Ok("back:x".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn notifications_reach_the_peer() {
|
||||
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
|
||||
let addr = start_server_with(move |session| {
|
||||
let session = session.clone();
|
||||
tokio::spawn(async move {
|
||||
session.notify::<Notice>("hello".into()).await.unwrap();
|
||||
});
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = Session::connect(&format!("ws://{addr}")).await.unwrap();
|
||||
client
|
||||
.on_notification::<Notice, _>(move |msg| {
|
||||
let tx = tx.clone();
|
||||
async move {
|
||||
tx.send(msg).unwrap();
|
||||
}
|
||||
})
|
||||
.await;
|
||||
client.start_receiver();
|
||||
|
||||
assert_eq!(
|
||||
tokio::time::timeout(WAIT, rx.recv()).await.unwrap(),
|
||||
Some("hello".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn close_fails_pending_requests_and_runs_on_close_once() {
|
||||
let closes = Arc::new(AtomicUsize::new(0));
|
||||
let (session_tx, mut session_rx) = tokio::sync::mpsc::unbounded_channel();
|
||||
let addr = start_server_with(move |session| {
|
||||
session_tx.send(session).unwrap();
|
||||
})
|
||||
.await;
|
||||
|
||||
let client = connect(addr).await;
|
||||
client
|
||||
.on_close({
|
||||
let closes = closes.clone();
|
||||
move || {
|
||||
closes.fetch_add(1, Ordering::SeqCst);
|
||||
async { Ok(()) }
|
||||
}
|
||||
})
|
||||
.await;
|
||||
|
||||
let pending = tokio::spawn({
|
||||
let client = client.clone();
|
||||
async move { client.request::<Silent>(()).await }
|
||||
});
|
||||
|
||||
let server_side = tokio::time::timeout(WAIT, session_rx.recv()).await.unwrap().unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
server_side.close().await.unwrap();
|
||||
|
||||
let result = tokio::time::timeout(WAIT, pending).await.unwrap().unwrap();
|
||||
assert!(matches!(result, Err(Error::ConnectionClosed)), "{result:?}");
|
||||
|
||||
tokio::time::timeout(WAIT, client.closed()).await.unwrap();
|
||||
client.close().await.unwrap();
|
||||
assert_eq!(closes.load(Ordering::SeqCst), 1);
|
||||
assert!(matches!(
|
||||
client.request::<Echo>("late".into()).await,
|
||||
Err(Error::ConnectionClosed)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_timeout() {
|
||||
let client = connect(start_server().await).await;
|
||||
|
||||
let result = client
|
||||
.request_timeout::<Silent>((), Duration::from_millis(100))
|
||||
.await;
|
||||
assert!(matches!(result, Err(Error::Timeout)), "{result:?}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oversized_message_closes_only_that_connection() {
|
||||
let addr = start_server().await;
|
||||
let mut ws = raw_client(addr).await;
|
||||
|
||||
let _ = ws.send(WsMessage::Text("x".repeat(2 << 20).into())).await;
|
||||
let closed = tokio::time::timeout(WAIT, async {
|
||||
loop {
|
||||
match ws.next().await {
|
||||
None | Some(Err(_)) | Some(Ok(WsMessage::Close(_))) => break,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert!(closed.is_ok(), "server kept the connection open");
|
||||
|
||||
let client = connect(addr).await;
|
||||
assert_eq!(client.request::<Echo>("alive".into()).await.unwrap(), Ok("alive".into()));
|
||||
}
|
||||
|
||||
/// Opens a TCP connection and completes a WebSocket handshake by hand.
|
||||
async fn raw_handshake(addr: SocketAddr) -> TcpStream {
|
||||
let mut tcp = TcpStream::connect(addr).await.unwrap();
|
||||
tcp.write_all(
|
||||
format!(
|
||||
"GET / HTTP/1.1\r\nHost: {addr}\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
|
||||
Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\n\r\n"
|
||||
)
|
||||
.as_bytes(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut response = Vec::new();
|
||||
while !response.ends_with(b"\r\n\r\n") {
|
||||
response.push(tcp.read_u8().await.unwrap());
|
||||
}
|
||||
assert!(response.starts_with(b"HTTP/1.1 101"));
|
||||
tcp
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn huge_frame_length_header_is_rejected_without_allocating() {
|
||||
let addr = start_server().await;
|
||||
let mut tcp = raw_handshake(addr).await;
|
||||
|
||||
// Masked text frame claiming a 2^62-byte payload. session-rs 0.1 tried to
|
||||
// allocate this up front.
|
||||
let mut frame = vec![0x81, 0x80 | 127];
|
||||
frame.extend_from_slice(&(1u64 << 62).to_be_bytes());
|
||||
frame.extend_from_slice(&[1, 2, 3, 4]);
|
||||
tcp.write_all(&frame).await.unwrap();
|
||||
|
||||
let mut buf = [0u8; 1024];
|
||||
let closed = tokio::time::timeout(WAIT, async {
|
||||
loop {
|
||||
match tcp.read(&mut buf).await {
|
||||
Ok(0) | Err(_) => break,
|
||||
Ok(_) => {}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await;
|
||||
assert!(closed.is_ok(), "server kept the connection open");
|
||||
|
||||
let client = connect(addr).await;
|
||||
assert_eq!(client.request::<Echo>("alive".into()).await.unwrap(), Ok("alive".into()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ping_timeout_closes_unresponsive_peer() {
|
||||
let closes = Arc::new(AtomicUsize::new(0));
|
||||
let addr = start_server_with({
|
||||
let closes = closes.clone();
|
||||
move |session| {
|
||||
let closes = closes.clone();
|
||||
tokio::spawn(async move {
|
||||
session
|
||||
.on_close(move || {
|
||||
closes.fetch_add(1, Ordering::SeqCst);
|
||||
async { Ok(()) }
|
||||
})
|
||||
.await;
|
||||
session.start_ping(Duration::from_millis(50), Duration::from_millis(100));
|
||||
});
|
||||
}
|
||||
})
|
||||
.await;
|
||||
|
||||
// A tungstenite client answers pings, so it must stay connected.
|
||||
let client = connect(addr).await;
|
||||
tokio::time::sleep(Duration::from_millis(400)).await;
|
||||
assert_eq!(client.request::<Echo>("alive".into()).await.unwrap(), Ok("alive".into()));
|
||||
assert_eq!(closes.load(Ordering::SeqCst), 0);
|
||||
|
||||
// A raw TCP peer never reads, so it never pongs.
|
||||
let _tcp = raw_handshake(addr).await;
|
||||
tokio::time::timeout(WAIT, async {
|
||||
while closes.load(Ordering::SeqCst) == 0 {
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("unresponsive peer was not closed");
|
||||
}
|
||||
|
||||
#[cfg(feature = "axum")]
|
||||
#[tokio::test]
|
||||
async fn axum_adapter_serves_sessions_next_to_http_routes() {
|
||||
use axum::{Router, extract::WebSocketUpgrade, routing::get};
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/",
|
||||
get(async |upgrade: WebSocketUpgrade| {
|
||||
upgrade.on_upgrade(async |socket| {
|
||||
let session = Session::from_axum(socket);
|
||||
register_handlers(&session).await;
|
||||
session.start_receiver();
|
||||
})
|
||||
}),
|
||||
)
|
||||
.route("/health", get(async || "ok"));
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move { axum::serve(listener, app).await });
|
||||
|
||||
let client = connect(addr).await;
|
||||
assert_eq!(client.request::<Echo>("hi".into()).await.unwrap(), Ok("hi".into()));
|
||||
assert_eq!(
|
||||
client.request::<AskBack>("y".into()).await.unwrap(),
|
||||
Ok("back:y".into())
|
||||
);
|
||||
|
||||
let mut tcp = TcpStream::connect(addr).await.unwrap();
|
||||
tcp.write_all(b"GET /health HTTP/1.1\r\nHost: x\r\nConnection: close\r\n\r\n")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut body = String::new();
|
||||
tcp.read_to_string(&mut body).await.unwrap();
|
||||
assert!(body.starts_with("HTTP/1.1 200") && body.ends_with("ok"), "{body}");
|
||||
}
|
||||
Reference in New Issue
Block a user