Files
session-rs/tests/protocol.rs
T
selimaj-devandclaude 05d5e53788 Release handlers on close to break session reference cycles
Handlers commonly capture a clone of their own Session, and the Session
owns its handlers, so a closed session was never freed. On shutdown,
take the on_close handler instead of cloning it and clear the request
and notification handlers after it runs.

Bump to 0.2.1.

Co-Authored-By: Claude Opus 5.5 <[email protected]>
2026-09-25 05:45:07 +02:00

493 lines
15 KiB
Rust

#![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}");
}
#[tokio::test]
async fn closing_releases_handlers_that_capture_the_session() {
let client = connect(start_server().await).await;
let captured = Arc::new(());
// Handlers typically hold a clone of their own session; that must not keep
// the session (and everything the handlers capture) alive after close.
client
.on_request::<Silent, _>({
let session = client.clone();
let captured = captured.clone();
move |_, ()| {
let _ = (&session, &captured);
async { Ok(()) }
}
})
.await;
client
.on_close({
let session = client.clone();
let captured = captured.clone();
move || {
let _ = (&session, &captured);
async { Ok(()) }
}
})
.await;
assert_eq!(Arc::strong_count(&captured), 3);
client.close().await.unwrap();
tokio::time::timeout(WAIT, async {
while Arc::strong_count(&captured) > 1 {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("handlers were not released after close");
}