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]>
493 lines
15 KiB
Rust
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");
|
|
}
|