Merge pull request '0.2.0: transport-agnostic protocol on tokio-tungstenite, plus axum adapter' (#1) from protocol-transport into master

Reviewed-on: #1
This commit was merged in pull request #1.
This commit is contained in:
2026-09-25 03:37:30 +00:00
16 changed files with 2230 additions and 804 deletions
Generated
+958 -30
View File
File diff suppressed because it is too large Load Diff
+31 -14
View File
@@ -1,21 +1,38 @@
[package] [package]
name = "session-rs" name = "session-rs"
version = "0.1.3" version = "0.2.0"
edition = "2024" edition = "2024"
description = "A lightweight async WebSocket protocol" description = "A lightweight async request/response and notification protocol over WebSockets"
license = "Apache-2.0" license = "Apache-2.0"
repository = "https://git.selimaj.dev/selimaj-dev/session-rs"
[features]
default = ["server", "client"]
# `SessionServer`: a standalone tokio-tungstenite WebSocket server.
server = ["dep:tokio-tungstenite", "tokio/net"]
# `Session::connect`: a tokio-tungstenite WebSocket client (ws:// only without a TLS feature).
client = ["dep:tokio-tungstenite", "tokio-tungstenite/connect"]
# `wss://` support for the client.
rustls = ["client", "tokio-tungstenite/rustls-tls-webpki-roots"]
native-tls = ["client", "tokio-tungstenite/native-tls"]
# `Session::from_axum`: run a session on an axum WebSocket upgrade.
axum = ["dep:axum"]
[dependencies] [dependencies]
base64 = "0.22.1" futures-util = { version = "0.3.34", default-features = false, features = ["sink", "std"] }
rand = "0.10.0" serde = { version = "1.0.228", features = ["derive"] }
serde = { version = "1.0.228", features = ["serde_derive"] }
serde_json = "1.0.149" serde_json = "1.0.149"
sha1 = "0.10.6" tokio = { version = "1.49.0", features = ["macros", "rt", "sync", "time"] }
tokio = { version = "1.49.0", features = [ tokio-tungstenite = { version = "0.30.0", default-features = false, features = ["handshake"], optional = true }
"io-util", axum = { version = "0.8.9", default-features = false, features = ["ws"], optional = true }
"macros",
"net", [dev-dependencies]
"rt", tokio = { version = "1.49.0", features = ["macros", "rt-multi-thread", "net", "time", "io-util"] }
"sync", axum = { version = "0.8.9", features = ["ws"] }
"time",
] } [[example]]
name = "axum"
required-features = ["axum"]
[package.metadata.docs.rs]
all-features = true
+77 -42
View File
@@ -6,20 +6,12 @@
## Introduction ## Introduction
This library provides **type-safe WebSocket communication** with a request-response and notification system built on top of a flexible protocol. `session-rs` is a small request/response + notification protocol that runs over WebSockets, with typed methods on both ends.
It ensures compile-time guarantees for message structure, reduces runtime errors, and simplifies building Rust client/server applications.
- **Dynamic Methods**: Each message includes a method enum for type safety. - **Typed methods**: requests, responses and errors are (de)serialized for you.
- **Typed Requests & Responses**: Automatic serialization and deserialization. - **Both directions**: either peer can send requests and notifications.
- **Optional Notifications**: Send asynchronous notifications across sessions. - **Transport-agnostic**: the protocol runs over any `Sink`/`Stream` of frames. Adapters ship for [tokio-tungstenite](https://docs.rs/tokio-tungstenite) and [axum](https://docs.rs/axum).
- **Bounded**: the built-in server limits message and frame sizes (1 MiB by default).
## Features
- Fully typed WebSocket sessions
- Type-safe request/response mechanism
- Optional typed notifications (Todo)
- Lightweight, minimal runtime overhead
- Async-first with Tokio support
## Installation ## Installation
@@ -27,9 +19,16 @@ It ensures compile-time guarantees for message structure, reduces runtime errors
cargo add session-rs cargo add session-rs
``` ```
--- | Feature | Default | Enables |
| --- | --- | --- |
| `server` | yes | `SessionServer`, a standalone WebSocket server |
| `client` | yes | `Session::connect` for `ws://` URLs |
| `rustls` / `native-tls` | no | `wss://` URLs in `Session::connect` |
| `axum` | no | `Session::from_axum` for axum WebSocket upgrades |
### **Basic Example (client)** ## Usage
Define a method once and share it between both peers:
```rust ```rust
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
@@ -41,44 +40,74 @@ impl Method for Data {
type Response = String; type Response = String;
type Error = String; type Error = String;
} }
let session = Session::connect("127.0.0.1:8080", "/").await?;
session.start_receiver();
session
.request::<Data>("Hello from client".to_string())
.await?;
``` ```
### **Basic Example (server)** ### Server
```rust ```rust
#[derive(Debug, Serialize, Deserialize)]
struct Data;
impl Method for Data {
const NAME: &'static str = "data";
type Request = String;
type Response = String;
type Error = String;
}
let server = SessionServer::bind("127.0.0.1:8080").await?; let server = SessionServer::bind("127.0.0.1:8080").await?;
server server
.session_loop(async |session, addr| { .session_loop(async |session, addr| {
// This will run on every new client // Runs for every new client. Register handlers here; the session
// starts reading once this returns, so no message arrives early.
session
.on_request::<Data, _>(async |_id, req| Ok(format!("echo: {req}")))
.await;
Ok(()) Ok(())
}).await; })
.await?;
``` ```
Use `SessionServer::with_config(ServerConfig { .. })` to change the size limits or the handshake timeout.
### Client
```rust
let session = Session::connect("ws://127.0.0.1:8080").await?;
session.start_receiver();
let reply = session.request::<Data>("Hello".to_string()).await?; // Ok("echo: Hello")
```
### axum
With the `axum` feature, a session can share a router with ordinary HTTP routes, such as a health check:
```rust
async fn ws(upgrade: WebSocketUpgrade) -> Response {
upgrade.max_message_size(1 << 20).on_upgrade(async |socket| {
let session = Session::from_axum(socket);
session.on_request::<Data, _>(async |_, req| Ok(req)).await;
session.start_receiver();
})
}
let app = Router::new()
.route("/", get(ws))
.route("/health", get(async || "ok"));
```
### Other transports
`Session::from_transport(sink, stream)` accepts any `Sink<Frame>` and `Stream<Item = Result<Frame, E>>`, and `Session::from_tungstenite` wraps an existing tokio-tungstenite stream (e.g. one accepted over TLS). The transport must answer pings itself.
### Semantics
- Incoming requests and notifications are handled one at a time, in arrival order. Responses to your own requests are delivered independently, so a handler can `request` from its peer.
- A request for an unknown method, or with data that doesn't deserialize, gets an error response instead of no reply.
- A handler that panics fails only its own request.
- `request` fails with `Error::ConnectionClosed` if the session closes first; `request_timeout` adds a deadline.
- `on_close` runs exactly once, whichever side closes. `start_ping(interval, timeout)` closes peers that stop answering pings.
## Protocol ## Protocol
Every message is a JSON text frame with a `type` tag.
#### Request #### Request
The request `id` is separated from the peer, and will increment only on it's requests. The `id` is chosen by the sender and increments per peer.
```json ```json
{ "type": "request", "id": 1, "method": "data", "data": "Hello from client" } { "type": "request", "id": 1, "method": "data", "data": "Hello from client" }
@@ -86,16 +115,22 @@ The request `id` is separated from the peer, and will increment only on it's req
#### Response #### Response
The response `id` **must** remain the same as the request. A response **must** carry the id of the request it answers.
```json ```json
{ "type": "response", "id": 1, "result": "Hello from server" } { "type": "response", "id": 1, "result": "Hello from server" }
``` ```
#### Notifications #### Error response
A notification is a method that doesn't need validation or output, it simply notifies a peer for a specific information
```json ```json
{ "type": "notification", "result": "Hello from server" } { "type": "errorresponse", "id": 1, "error": "Invalid data" }
```
#### Notification
A notification is fire-and-forget and gets no response.
```json
{ "type": "notification", "method": "data", "data": "Hello from server" }
``` ```
+41
View File
@@ -0,0 +1,41 @@
//! The same server as `examples/server.rs`, mounted in an axum router next to
//! ordinary HTTP routes.
use axum::{Router, extract::WebSocketUpgrade, response::Response, routing::get};
use serde::{Deserialize, Serialize};
use session_rs::{Method, Session};
#[derive(Debug, Serialize, Deserialize)]
struct Data;
impl Method for Data {
const NAME: &'static str = "data";
type Request = String;
type Response = String;
type Error = String;
}
async fn ws(upgrade: WebSocketUpgrade) -> Response {
upgrade.max_message_size(1 << 20).on_upgrade(async |socket| {
let session = Session::from_axum(socket);
session
.on_request::<Data, _>(async |_, req| {
println!("Msg from client: {req}");
Ok(format!("Hello from axum, you said {req:?}"))
})
.await;
session.start_receiver();
})
}
#[tokio::main]
async fn main() -> std::io::Result<()> {
let app = Router::new()
.route("/", get(ws))
.route("/health", get(async || "ok"));
let listener = tokio::net::TcpListener::bind("127.0.0.1:8080").await?;
axum::serve(listener, app).await
}
+2 -4
View File
@@ -1,5 +1,5 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use session_rs::{Method, session::Session}; use session_rs::{Method, Session};
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
struct Data; struct Data;
@@ -13,7 +13,7 @@ impl Method for Data {
#[tokio::main(flavor = "current_thread")] #[tokio::main(flavor = "current_thread")]
async fn main() -> session_rs::Result<()> { async fn main() -> session_rs::Result<()> {
let session = Session::connect("127.0.0.1:8080", "/").await?; let session = Session::connect("ws://127.0.0.1:8080").await?;
session.start_receiver(); session.start_receiver();
@@ -31,8 +31,6 @@ async fn main() -> session_rs::Result<()> {
.await? .await?
); );
tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await;
session.close().await?; session.close().await?;
Ok(()) Ok(())
} }
+39
View File
@@ -0,0 +1,39 @@
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use crate::{Frame, Session};
fn to_ws(frame: Frame) -> Message {
match frame {
Frame::Text(text) => Message::Text(text.into()),
Frame::Binary(data) => Message::Binary(data.into()),
Frame::Ping(data) => Message::Ping(data.into()),
Frame::Pong(data) => Message::Pong(data.into()),
Frame::Close => Message::Close(None),
}
}
fn from_ws(msg: Message) -> Frame {
match msg {
Message::Text(text) => Frame::Text(text.as_str().to_owned()),
Message::Binary(data) => Frame::Binary(data.to_vec()),
Message::Ping(data) => Frame::Ping(data.to_vec()),
Message::Pong(data) => Frame::Pong(data.to_vec()),
Message::Close(_) => Frame::Close,
}
}
impl Session {
/// Run a session over an axum WebSocket upgrade.
///
/// Message size limits are set on the upgrade, e.g.
/// `ws.max_message_size(1 << 20).on_upgrade(...)`.
pub fn from_axum(socket: WebSocket) -> Self {
let (sink, stream) = socket.split();
Session::from_transport(
sink.with(|frame| async move { Ok::<_, axum::Error>(to_ws(frame)) }),
stream.map(|msg| msg.map(from_ws)),
)
}
}
+13
View File
@@ -0,0 +1,13 @@
use crate::{Error, Session};
impl Session {
/// Connect to a `ws://` (or, with the `rustls`/`native-tls` feature,
/// `wss://`) URL. Call [`Session::start_receiver`] after registering handlers.
pub async fn connect(url: &str) -> crate::Result<Self> {
let (ws, _) = tokio_tungstenite::connect_async(url)
.await
.map_err(|e| Error::Transport(Box::new(e)))?;
Ok(Self::from_tungstenite(ws))
}
}
+45 -13
View File
@@ -1,16 +1,35 @@
use std::{pin::Pin, sync::Arc}; //! A small request/response + notification protocol over WebSockets.
//!
//! The protocol layer ([`Session`]) is transport-agnostic: it runs over any
//! [`Sink`](futures_util::Sink)/[`Stream`](futures_util::Stream) pair of
//! [`Frame`]s. Adapters are provided for tokio-tungstenite (`server`/`client`
//! features) and axum (`axum` feature).
use std::pin::Pin;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
#[cfg(feature = "axum")]
mod axum;
#[cfg(feature = "client")]
mod client;
#[cfg(feature = "server")]
pub mod server; pub mod server;
pub mod session; pub mod session;
pub mod ws; pub mod transport;
#[cfg(any(feature = "server", feature = "client"))]
mod tungstenite;
pub use session::{Message, Session};
pub use transport::Frame;
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> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
Pin<Box<dyn Future<Output = T> + Send + 'a>>; pub type BoxError = Box<dyn std::error::Error + Send + Sync>;
pub type MethodHandler = Arc<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>;
/// A named, typed RPC. Implement this on a marker type and use it with
/// [`Session::request`], [`Session::on_request`], [`Session::notify`] and
/// [`Session::on_notification`].
pub trait Method { pub trait Method {
const NAME: &'static str; const NAME: &'static str;
type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync; type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync;
@@ -18,6 +37,7 @@ pub trait Method {
type Error: Serialize + for<'de> Deserialize<'de>; type Error: Serialize + for<'de> Deserialize<'de>;
} }
/// Untyped method used internally for raw JSON values.
pub struct GenericMethod; pub struct GenericMethod;
impl Method for GenericMethod { impl Method for GenericMethod {
@@ -29,18 +49,30 @@ impl Method for GenericMethod {
#[derive(Debug)] #[derive(Debug)]
pub enum Error { pub enum Error {
WebSocket(ws::Error), /// The underlying transport (WebSocket) failed.
Transport(BoxError),
Json(serde_json::Error), Json(serde_json::Error),
Io(std::io::Error), Io(std::io::Error),
RecvError(tokio::sync::broadcast::error::RecvError), /// The session is closed; no more messages can be sent or received.
ConnectionClosed,
/// A request or handshake did not complete in time.
Timeout,
} }
impl From<ws::Error> for Error { impl std::fmt::Display for Error {
fn from(value: ws::Error) -> Self { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
Self::WebSocket(value) match self {
Error::Transport(e) => write!(f, "transport error: {e}"),
Error::Json(e) => write!(f, "json error: {e}"),
Error::Io(e) => write!(f, "io error: {e}"),
Error::ConnectionClosed => write!(f, "connection closed"),
Error::Timeout => write!(f, "timed out"),
}
} }
} }
impl std::error::Error for Error {}
impl From<std::io::Error> for Error { impl From<std::io::Error> for Error {
fn from(value: std::io::Error) -> Self { fn from(value: std::io::Error) -> Self {
Self::Io(value) Self::Io(value)
@@ -53,8 +85,8 @@ impl From<serde_json::Error> for Error {
} }
} }
impl From<tokio::sync::broadcast::error::RecvError> for Error { impl From<tokio::time::error::Elapsed> for Error {
fn from(value: tokio::sync::broadcast::error::RecvError) -> Self { fn from(_: tokio::time::error::Elapsed) -> Self {
Self::RecvError(value) Self::Timeout
} }
} }
+95 -33
View File
@@ -1,62 +1,124 @@
use std::{net::SocketAddr, sync::Arc}; use std::{net::SocketAddr, sync::Arc, time::Duration};
use tokio::{net::TcpListener, time::timeout}; use tokio::net::{TcpListener, TcpStream, ToSocketAddrs};
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
use crate::{session::Session, ws::WebSocket}; use crate::{Error, session::Session};
/// Limits applied to every accepted connection.
#[derive(Debug, Clone)]
pub struct ServerConfig {
/// Largest message (after reassembling fragments) a peer may send.
pub max_message_size: usize,
/// Largest single frame a peer may send.
pub max_frame_size: usize,
/// How long a client gets to complete the WebSocket handshake.
pub handshake_timeout: Duration,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
max_message_size: 1 << 20,
max_frame_size: 1 << 20,
handshake_timeout: Duration::from_secs(5),
}
}
}
/// A standalone WebSocket server that hands each connection to a callback as a
/// [`Session`].
pub struct SessionServer { pub struct SessionServer {
listener: TcpListener, listener: TcpListener,
config: ServerConfig,
} }
impl SessionServer { impl SessionServer {
pub async fn bind(addr: &str) -> crate::Result<Self> { pub async fn bind(addr: impl ToSocketAddrs) -> crate::Result<Self> {
Ok(Self { Ok(Self::from_listener(TcpListener::bind(addr).await?))
listener: TcpListener::bind(addr).await?,
})
} }
pub fn from_listener(listener: TcpListener) -> Self {
Self {
listener,
config: ServerConfig::default(),
}
}
pub fn with_config(mut self, config: ServerConfig) -> Self {
self.config = config;
self
}
pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.listener.local_addr()
}
/// Accept one connection. The receiver is not started: register handlers,
/// then call [`Session::start_receiver`].
pub async fn accept(&self) -> crate::Result<(Session, SocketAddr)> { pub async fn accept(&self) -> crate::Result<(Session, SocketAddr)> {
let (stream, addr) = self.listener.accept().await?; let (stream, addr) = self.listener.accept().await?;
Ok((handshake(stream, &self.config).await?, addr))
let ws = WebSocket::handshake(stream).await?;
Ok((Session::from_ws(ws), addr))
} }
/// Accept connections forever, running `on_conn` for each one.
///
/// `on_conn` should register handlers and return; the receiver starts once
/// it does, so no message is processed before its handler exists. If it
/// returns an error the connection is closed.
pub async fn session_loop<F, Fut>(&self, on_conn: F) -> crate::Result<()> pub async fn session_loop<F, Fut>(&self, on_conn: F) -> crate::Result<()>
where where
F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static, F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static,
Fut: Future<Output = crate::Result<()>> + Send + 'static, Fut: Future<Output = crate::Result<()>> + Send + 'static,
{ {
let conn_handler = Arc::new(on_conn); let on_conn = Arc::new(on_conn);
loop { loop {
let (stream, addr) = self.listener.accept().await?; let (stream, addr) = match self.listener.accept().await {
let conn_handler = conn_handler.clone(); Ok(conn) => conn,
Err(e) => {
// e.g. out of file descriptors: back off instead of exiting.
eprintln!("Accept failed: {e}");
tokio::time::sleep(Duration::from_millis(100)).await;
continue;
}
};
let on_conn = on_conn.clone();
let config = self.config.clone();
tokio::spawn(async move { tokio::spawn(async move {
match timeout( let session = match handshake(stream, &config).await {
tokio::time::Duration::from_secs(5), Ok(session) => session,
WebSocket::handshake(stream), Err(e) => {
) eprintln!("Handshake failed from {addr}: {e}");
.await return;
{ }
Ok(Ok(ws)) => { };
let session = Session::from_ws(ws);
session.start_receiver();
if let Err(e) = conn_handler(session, addr).await { if let Err(e) = on_conn(session.clone(), addr).await {
eprintln!("Connection error: {:?}", e); eprintln!("Connection error from {addr}: {e}");
} let _ = session.close().await;
} return;
Ok(Err(e)) => {
eprintln!("Handshake failed from {}: {:?}", addr, e);
}
Err(_) => {
eprintln!("Handshake failed from {}: Handshake Timeout", addr);
}
} }
session.start_receiver();
}); });
} }
} }
} }
async fn handshake(stream: TcpStream, config: &ServerConfig) -> crate::Result<Session> {
let ws_config = WebSocketConfig::default()
.max_message_size(Some(config.max_message_size))
.max_frame_size(Some(config.max_frame_size));
let ws = tokio::time::timeout(
config.handshake_timeout,
tokio_tungstenite::accept_async_with_config(stream, Some(ws_config)),
)
.await?
.map_err(|e| Error::Transport(Box::new(e)))?;
Ok(Session::from_tungstenite(ws))
}
+394 -161
View File
@@ -1,14 +1,25 @@
use std::collections::HashMap;
use std::hash::Hash; use std::hash::Hash;
use std::{collections::HashMap, sync::Arc}; use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tokio::sync::Mutex; use serde_json::Value;
use tokio::sync::broadcast; use tokio::sync::{mpsc, oneshot, watch};
use tokio::time::timeout;
use crate::BoxFuture; use crate::transport::{self, BoxSink, BoxStream, Frame};
use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket}; use crate::{BoxError, BoxFuture, Error, GenericMethod, Method};
/// A protocol message, serialized as JSON in a text frame.
///
/// ```json
/// { "type": "request", "id": 1, "method": "data", "data": "..." }
/// { "type": "response", "id": 1, "result": "..." }
/// { "type": "errorresponse", "id": 1, "error": "..." }
/// { "type": "notification", "method": "data", "data": "..." }
/// ```
#[derive(Debug, Serialize, Deserialize)] #[derive(Debug, Serialize, Deserialize)]
#[serde(rename_all = "lowercase", tag = "type")] #[serde(rename_all = "lowercase", tag = "type")]
pub enum Message<M: Method> { pub enum Message<M: Method> {
@@ -31,126 +42,246 @@ pub enum Message<M: Method> {
}, },
} }
type RequestHandler = Arc<dyn Fn(u32, Value) -> BoxFuture<'static, Result<Value, Value>> + Send + Sync>;
type NotificationHandler = Arc<dyn Fn(Value) -> BoxFuture<'static, ()> + Send + Sync>;
type CloseHandler = Arc<dyn Fn() -> BoxFuture<'static, Result<(), String>> + Send + Sync>;
/// Frames queued for the writer task before `send` applies backpressure.
const OUTGOING_BUFFER: usize = 256;
/// Requests/notifications queued for handlers before the reader applies backpressure.
const INCOMING_BUFFER: usize = 64;
static NEXT_SESSION_ID: AtomicU64 = AtomicU64::new(1);
/// One peer connection speaking the session protocol.
///
/// Cloning is cheap and every clone refers to the same connection; equality and
/// hashing are by connection, so sessions can be stored in sets and maps.
///
/// Incoming requests and notifications are handled one at a time, in arrival
/// order. Responses to this side's own requests are delivered independently,
/// so a handler may itself `request` from the peer.
#[derive(Clone)]
pub struct Session { pub struct Session {
pub ws: WebSocket, inner: Arc<Inner>,
id: Arc<Mutex<u32>>, }
methods: Arc<Mutex<HashMap<String, MethodHandler>>>,
on_close_fn: struct Inner {
Arc<Mutex<Option<Box<dyn Fn() -> BoxFuture<'static, Result<(), String>> + Send + Sync>>>>, id: u64,
tx: broadcast::Sender<(u32, bool, serde_json::Value)>, next_request_id: AtomicU32,
pong_tx: broadcast::Sender<()>, outgoing: mpsc::Sender<Frame>,
/// The read half, until `start_receiver` takes it.
stream: Mutex<Option<BoxStream>>,
requests: Mutex<HashMap<String, RequestHandler>>,
notifications: Mutex<HashMap<String, NotificationHandler>>,
pending: Mutex<HashMap<u32, oneshot::Sender<Result<Value, Value>>>>,
on_close: Mutex<Option<CloseHandler>>,
closed: AtomicBool,
closed_tx: watch::Sender<bool>,
/// Bumped on every pong, for `start_ping`.
pongs: watch::Sender<u64>,
}
enum Incoming {
Request { id: u32, method: String, data: Value },
Notification { method: String, data: Value },
} }
impl Session { impl Session {
pub fn clone(&self) -> Self { /// Build a session over any frame sink/stream pair.
Self { ///
ws: self.ws.clone(), /// The transport is expected to answer pings itself (tungstenite and axum
id: self.id.clone(), /// both do). Nothing is read until [`Session::start_receiver`] is called,
methods: self.methods.clone(), /// so handlers can be registered first.
on_close_fn: self.on_close_fn.clone(), pub fn from_transport<Si, St, SiE, StE>(sink: Si, stream: St) -> Self
tx: self.tx.clone(), where
pong_tx: self.pong_tx.clone(), Si: Sink<Frame, Error = SiE> + Send + 'static,
} SiE: Into<BoxError> + 'static,
St: Stream<Item = Result<Frame, StE>> + Send + 'static,
StE: Into<BoxError> + 'static,
{
let (sink, stream) = transport::boxed(sink, stream);
let (outgoing, outgoing_rx) = mpsc::channel(OUTGOING_BUFFER);
let (closed_tx, _) = watch::channel(false);
let session = Self {
inner: Arc::new(Inner {
id: NEXT_SESSION_ID.fetch_add(1, Ordering::Relaxed),
next_request_id: AtomicU32::new(0),
outgoing,
stream: Mutex::new(Some(stream)),
requests: Mutex::new(HashMap::new()),
notifications: Mutex::new(HashMap::new()),
pending: Mutex::new(HashMap::new()),
on_close: Mutex::new(None),
closed: AtomicBool::new(false),
closed_tx,
pongs: watch::channel(0).0,
}),
};
tokio::spawn(write_loop(
outgoing_rx,
sink,
session.inner.closed_tx.subscribe(),
Arc::downgrade(&session.inner),
));
session
} }
}
impl Session { /// A process-unique id for this connection.
pub fn from_ws(ws: WebSocket) -> Self { pub fn id(&self) -> u64 {
let (tx, _) = broadcast::channel(8192); self.inner.id
let (pong_tx, _) = broadcast::channel(16); }
Self { pub fn is_closed(&self) -> bool {
ws, self.inner.closed.load(Ordering::SeqCst)
id: Arc::new(Mutex::new(0)),
methods: Arc::new(Mutex::new(HashMap::new())),
on_close_fn: Arc::new(Mutex::new(None)),
tx,
pong_tx,
}
} }
pub async fn connect(addr: &str, path: &str) -> crate::Result<Self> { /// Resolves once the session is closed, from either side.
Ok(Self::from_ws(WebSocket::connect(addr, path).await?)) pub async fn closed(&self) {
wait_closed(&mut self.inner.closed_tx.subscribe()).await;
} }
} }
impl Session { impl Session {
/// Start reading from the peer. Calling it again has no effect.
pub fn start_receiver(&self) { pub fn start_receiver(&self) {
let Some(stream) = self.inner.stream.lock().unwrap().take() else {
return;
};
let (dispatch_tx, dispatch_rx) = mpsc::channel(INCOMING_BUFFER);
tokio::spawn(self.clone().dispatch_loop(dispatch_rx));
tokio::spawn(self.clone().read_loop(stream, dispatch_tx));
}
/// Ping the peer every `interval` and close the session if no pong arrives
/// within `timeout`.
pub fn start_ping(&self, interval: Duration, timeout: Duration) {
let s = self.clone(); let s = self.clone();
tokio::spawn(async move { tokio::spawn(async move {
let mut closed = s.inner.closed_tx.subscribe();
let mut pongs = s.inner.pongs.subscribe();
loop { loop {
match s.ws.read().await { tokio::select! {
Ok(crate::ws::Frame::Text(text)) => { _ = tokio::time::sleep(interval) => {}
let Ok(msg) = serde_json::from_str::<Message<GenericMethod>>(&text) else { _ = wait_closed(&mut closed) => return,
continue; }
};
match msg { pongs.mark_unchanged();
Message::Request { id, method, data } => {
let handler = {
let methods = s.methods.lock().await;
methods.get(&method).cloned()
};
if let Some(m) = handler { if s.inner.outgoing.send(Frame::Ping(Vec::new())).await.is_err() {
if let Some((err, res)) = (m)(id, data).await { break;
if err { }
s.respond_error(id, res)
.await match tokio::time::timeout(timeout, pongs.changed()).await {
.expect("Failed to respond"); Ok(Ok(())) => {}
} else { _ => break,
s.respond(id, res).await.expect("Failed to respond"); }
} }
}
} s.shutdown().await;
} });
Message::Response { id, result } => { }
s.tx.send((id, false, result)).unwrap();
} async fn read_loop(self, mut stream: BoxStream, dispatch: mpsc::Sender<Incoming>) {
Message::ErrorResponse { id, error } => { let mut closed = self.inner.closed_tx.subscribe();
s.tx.send((id, true, error)).unwrap();
} loop {
_ => {} let frame = tokio::select! {
} frame = stream.next() => frame,
} _ = wait_closed(&mut closed) => break,
Ok(crate::ws::Frame::Pong) => { };
let _ = s.pong_tx.send(());
} match frame {
Ok(_) => {} Some(Ok(Frame::Text(text))) => {
Err(_) => { if !self.handle_text(&text, &dispatch).await {
s.trigger_close().await;
break; break;
} }
} }
} Some(Ok(Frame::Pong(_))) => {
}); self.inner.pongs.send_modify(|n| *n = n.wrapping_add(1));
}
pub fn start_ping(&self, interval: tokio::time::Duration, timeout_dur: tokio::time::Duration) {
let s = self.clone();
tokio::spawn(async move {
let mut pong_rx = s.pong_tx.subscribe();
loop {
tokio::time::sleep(interval).await;
if s.ws.send_ping().await.is_err() {
s.trigger_close().await;
break;
}
let result = timeout(timeout_dur, pong_rx.recv()).await;
if result.is_err() {
// timeout expired
let _ = s.close().await;
s.trigger_close().await;
break;
} }
Some(Ok(Frame::Ping(_) | Frame::Binary(_))) => {}
Some(Ok(Frame::Close)) | Some(Err(_)) | None => break,
} }
}); }
self.shutdown().await;
} }
/// Returns false once the dispatcher is gone.
async fn handle_text(&self, text: &str, dispatch: &mpsc::Sender<Incoming>) -> bool {
let Ok(msg) = serde_json::from_str::<Message<GenericMethod>>(text) else {
return true;
};
let incoming = match msg {
Message::Request { id, method, data } => Incoming::Request { id, method, data },
Message::Notification { method, data } => Incoming::Notification { method, data },
Message::Response { id, result } => {
self.complete(id, Ok(result));
return true;
}
Message::ErrorResponse { id, error } => {
self.complete(id, Err(error));
return true;
}
};
dispatch.send(incoming).await.is_ok()
}
fn complete(&self, id: u32, result: Result<Value, Value>) {
if let Some(tx) = self.inner.pending.lock().unwrap().remove(&id) {
let _ = tx.send(result);
}
}
async fn dispatch_loop(self, mut rx: mpsc::Receiver<Incoming>) {
while let Some(incoming) = rx.recv().await {
match incoming {
Incoming::Request { id, method, data } => {
let handler = self.inner.requests.lock().unwrap().get(&method).cloned();
let result = match handler {
// Spawned so a panicking handler fails this request, not the session.
Some(handler) => tokio::spawn(handler(id, data))
.await
.unwrap_or_else(|_| Err(Value::from("Handler panicked"))),
None => Err(Value::from(format!("Unknown method: {method}"))),
};
let reply = match result {
Ok(v) => self.respond(id, v).await,
Err(e) => self.respond_error(id, e).await,
};
if reply.is_err() {
break;
}
}
Incoming::Notification { method, data } => {
let handler = self.inner.notifications.lock().unwrap().get(&method).cloned();
if let Some(handler) = handler {
let _ = tokio::spawn(handler(data)).await;
}
}
}
}
}
}
impl Session {
/// Handle requests for `M`. Replaces any previous handler for the method.
///
/// Requests whose data does not deserialize into `M::Request` are answered
/// with an error response.
pub async fn on_request< pub async fn on_request<
M: Method, M: Method,
Fut: Future<Output = Result<M::Response, M::Error>> + Send + 'static, Fut: Future<Output = Result<M::Response, M::Error>> + Send + 'static,
@@ -158,57 +289,89 @@ impl Session {
&self, &self,
handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static, handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static,
) { ) {
let handler = Arc::new(handler); let handler: RequestHandler = Arc::new(move |id, value| {
let fut = serde_json::from_value::<M::Request>(value).map(|req| handler(id, req));
self.methods.lock().await.insert( Box::pin(async move {
M::NAME.to_string(), match fut {
Arc::new(move |id, value| { Err(e) => Err(Value::from(format!("Invalid request data: {e}"))),
let handler = Arc::clone(&handler); Ok(fut) => match fut.await {
Ok(res) => serde_json::to_value(res)
.map_err(|e| Value::from(format!("Failed to serialize response: {e}"))),
Err(err) => Err(serde_json::to_value(err).unwrap_or_else(|e| {
Value::from(format!("Failed to serialize error: {e}"))
})),
},
}
})
});
Box::pin(async move { self.inner
Some( .requests
match handler(id, serde_json::from_value(value).ok()?).await { .lock()
Ok(v) => (false, serde_json::to_value(v).ok()?), .unwrap()
Err(v) => (true, serde_json::to_value(v).ok()?), .insert(M::NAME.to_string(), handler);
},
)
})
}),
);
} }
/// Handle notifications for `M`. Notifications with invalid data are dropped.
pub async fn on_notification<M: Method, Fut: Future<Output = ()> + Send + 'static>(
&self,
handler: impl Fn(M::Request) -> Fut + Send + Sync + 'static,
) {
let handler: NotificationHandler = Arc::new(move |value| {
let fut = serde_json::from_value::<M::Request>(value).map(&handler);
Box::pin(async move {
if let Ok(fut) = fut {
fut.await;
}
})
});
self.inner
.notifications
.lock()
.unwrap()
.insert(M::NAME.to_string(), handler);
}
/// Run `handler` once when the session closes, from either side.
pub async fn on_close<Fut>(&self, handler: impl Fn() -> Fut + Send + Sync + 'static) pub async fn on_close<Fut>(&self, handler: impl Fn() -> Fut + Send + Sync + 'static)
where where
Fut: Future<Output = Result<(), String>> + Send + 'static, Fut: Future<Output = Result<(), String>> + Send + 'static,
{ {
let handler = Arc::new(handler); *self.inner.on_close.lock().unwrap() = Some(Arc::new(move || Box::pin(handler())));
*self.on_close_fn.lock().await = Some(Box::new(move || {
let handler = handler.clone();
Box::pin(async move { handler().await })
}));
} }
} }
impl Session { impl Session {
pub async fn send<M: Method>(&self, data: &Message<M>) -> crate::Result<()> { pub async fn send<M: Method>(&self, msg: &Message<M>) -> crate::Result<()> {
self.ws if self.is_closed() {
.send_text_payload(&serde_json::to_vec(&data)?) return Err(Error::ConnectionClosed);
.await?; }
Ok(())
}
pub async fn use_id(&self) -> u32 { let text = serde_json::to_string(msg)?;
let mut id = self.id.lock().await;
*id += 1; self.inner
*id .outgoing
.send(Frame::Text(text))
.await
.map_err(|_| Error::ConnectionClosed)
} }
/// Send a request and wait for the peer's response.
///
/// Fails with [`Error::ConnectionClosed`] if the session closes first.
pub async fn request<M: Method>( pub async fn request<M: Method>(
&self, &self,
req: M::Request, req: M::Request,
) -> crate::Result<std::result::Result<M::Response, M::Error>> { ) -> crate::Result<std::result::Result<M::Response, M::Error>> {
let id = self.use_id().await; let id = self.inner.next_request_id.fetch_add(1, Ordering::Relaxed).wrapping_add(1);
let (tx, rx) = oneshot::channel();
// Registered before sending so a fast response can't be missed.
self.inner.pending.lock().unwrap().insert(id, tx);
let _guard = PendingGuard { inner: &self.inner, id };
self.send::<M>(&Message::Request { self.send::<M>(&Message::Request {
id, id,
@@ -217,30 +380,27 @@ impl Session {
}) })
.await?; .await?;
let mut rx = self.tx.subscribe(); Ok(match rx.await.map_err(|_| Error::ConnectionClosed)? {
Ok(v) => Ok(serde_json::from_value(v)?),
loop { Err(e) => Err(serde_json::from_value(e)?),
let r = rx.recv().await?;
if r.0 == id {
break Ok(if r.1 {
Err(serde_json::from_value(r.2)?)
} else {
Ok(serde_json::from_value(r.2)?)
});
}
}
}
pub async fn respond(&self, to: u32, val: serde_json::Value) -> crate::Result<()> {
self.send::<GenericMethod>(&Message::Response {
id: to,
result: val,
}) })
.await
} }
pub async fn respond_error(&self, to: u32, val: serde_json::Value) -> crate::Result<()> { /// [`Session::request`] with a deadline; fails with [`Error::Timeout`].
pub async fn request_timeout<M: Method>(
&self,
req: M::Request,
timeout: Duration,
) -> crate::Result<std::result::Result<M::Response, M::Error>> {
tokio::time::timeout(timeout, self.request::<M>(req)).await?
}
pub async fn respond(&self, to: u32, val: Value) -> crate::Result<()> {
self.send::<GenericMethod>(&Message::Response { id: to, result: val })
.await
}
pub async fn respond_error(&self, to: u32, val: Value) -> crate::Result<()> {
self.send::<GenericMethod>(&Message::ErrorResponse { id: to, error: val }) self.send::<GenericMethod>(&Message::ErrorResponse { id: to, error: val })
.await .await
} }
@@ -253,28 +413,101 @@ impl Session {
.await .await
} }
async fn trigger_close(&self) { /// Close the connection. Queued messages are flushed first.
if let Some(handler) = self.on_close_fn.lock().await.as_ref() { pub async fn close(&self) -> crate::Result<()> {
self.shutdown().await;
Ok(())
}
/// Marks the session closed, fails pending requests and runs `on_close`,
/// exactly once.
async fn shutdown(&self) {
if self.inner.closed.swap(true, Ordering::SeqCst) {
return;
}
self.inner.closed_tx.send_replace(true);
// Dropping the senders fails every waiting `request`.
drop(std::mem::take(&mut *self.inner.pending.lock().unwrap()));
let handler = self.inner.on_close.lock().unwrap().clone();
if let Some(handler) = handler {
let _ = handler().await; let _ = handler().await;
} }
} }
}
pub async fn close(&self) -> crate::Result<()> { /// Removes a pending request if its `request` future is dropped early.
let res = self.ws.close().await; struct PendingGuard<'a> {
self.trigger_close().await; inner: &'a Inner,
Ok(res?) id: u32,
}
impl Drop for PendingGuard<'_> {
fn drop(&mut self) {
self.inner.pending.lock().unwrap().remove(&self.id);
}
}
/// Waits for the closed flag. The `watch::Ref` is dropped inside, so callers
/// can use this in `select!` without holding a non-`Send` guard.
async fn wait_closed(closed: &mut watch::Receiver<bool>) {
let _ = closed.wait_for(|c| *c).await;
}
async fn write_loop(
mut rx: mpsc::Receiver<Frame>,
mut sink: BoxSink,
mut closed: watch::Receiver<bool>,
session: std::sync::Weak<Inner>,
) {
loop {
tokio::select! {
biased;
frame = rx.recv() => {
let Some(frame) = frame else { break };
if sink.send(frame).await.is_err() {
break;
}
}
_ = wait_closed(&mut closed) => {
while let Ok(frame) = rx.try_recv() {
if sink.send(frame).await.is_err() {
break;
}
}
let _ = sink.send(Frame::Close).await;
break;
}
}
}
let _ = sink.close().await;
if let Some(inner) = session.upgrade() {
Session { inner }.shutdown().await;
}
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session")
.field("id", &self.inner.id)
.field("closed", &self.is_closed())
.finish()
} }
} }
impl Hash for Session { impl Hash for Session {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) { fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.ws.id.hash(state); self.inner.id.hash(state);
} }
} }
impl PartialEq for Session { impl PartialEq for Session {
fn eq(&self, other: &Self) -> bool { fn eq(&self, other: &Self) -> bool {
self.ws.id == other.ws.id self.inner.id == other.inner.id
} }
} }
+36
View File
@@ -0,0 +1,36 @@
//! The transport boundary between the protocol and a WebSocket implementation.
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use crate::BoxError;
/// A WebSocket-level frame as seen by the protocol layer.
///
/// Adapters translate their library's message type to and from this. Only
/// `Text` carries protocol messages; `Ping`/`Pong` drive liveness checks and
/// `Close` ends the session.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Frame {
Text(String),
Binary(Vec<u8>),
Ping(Vec<u8>),
Pong(Vec<u8>),
Close,
}
pub(crate) type BoxSink = std::pin::Pin<Box<dyn Sink<Frame, Error = BoxError> + Send>>;
pub(crate) type BoxStream =
std::pin::Pin<Box<dyn Stream<Item = Result<Frame, BoxError>> + Send>>;
pub(crate) fn boxed<Si, St, SiE, StE>(sink: Si, stream: St) -> (BoxSink, BoxStream)
where
Si: Sink<Frame, Error = SiE> + Send + 'static,
SiE: Into<BoxError> + 'static,
St: Stream<Item = Result<Frame, StE>> + Send + 'static,
StE: Into<BoxError> + 'static,
{
(
Box::pin(sink.sink_map_err(Into::into)),
Box::pin(stream.map(|r| r.map_err(Into::into))),
)
}
+47
View File
@@ -0,0 +1,47 @@
use futures_util::{SinkExt, StreamExt};
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_tungstenite::{WebSocketStream, tungstenite};
use crate::{Frame, Session};
fn to_ws(frame: Frame) -> tungstenite::Message {
match frame {
Frame::Text(text) => tungstenite::Message::Text(text.into()),
Frame::Binary(data) => tungstenite::Message::Binary(data.into()),
Frame::Ping(data) => tungstenite::Message::Ping(data.into()),
Frame::Pong(data) => tungstenite::Message::Pong(data.into()),
Frame::Close => tungstenite::Message::Close(None),
}
}
fn from_ws(msg: tungstenite::Message) -> Option<Frame> {
Some(match msg {
tungstenite::Message::Text(text) => Frame::Text(text.as_str().to_owned()),
tungstenite::Message::Binary(data) => Frame::Binary(data.to_vec()),
tungstenite::Message::Ping(data) => Frame::Ping(data.to_vec()),
tungstenite::Message::Pong(data) => Frame::Pong(data.to_vec()),
tungstenite::Message::Close(_) => Frame::Close,
tungstenite::Message::Frame(_) => return None,
})
}
impl Session {
/// Run a session over an established tokio-tungstenite WebSocket, e.g. one
/// accepted with a custom TLS acceptor or handshake callback.
pub fn from_tungstenite<S>(ws: WebSocketStream<S>) -> Self
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let (sink, stream) = ws.split();
Session::from_transport(
sink.with(|frame| async move { Ok::<_, tungstenite::Error>(to_ws(frame)) }),
stream.filter_map(|msg| async move {
match msg {
Ok(msg) => from_ws(msg).map(Ok),
Err(e) => Some(Err(e)),
}
}),
)
}
}
-31
View File
@@ -1,31 +0,0 @@
use std::string::FromUtf8Error;
pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug)]
pub enum Error {
Io(std::io::Error),
InvalidFrame(String),
HandshakeFailed(String),
Utf8(FromUtf8Error),
ConnectionClosed,
Elapsed,
}
impl From<std::io::Error> for Error {
fn from(value: std::io::Error) -> Self {
Self::Io(value)
}
}
impl From<FromUtf8Error> for Error {
fn from(value: FromUtf8Error) -> Self {
Self::Utf8(value)
}
}
impl From<tokio::time::error::Elapsed> for Error {
fn from(_: tokio::time::error::Elapsed) -> Self {
Self::Elapsed
}
}
-225
View File
@@ -1,225 +0,0 @@
use base64::Engine;
use base64::engine::general_purpose::STANDARD as Base64;
use sha1::{Digest, Sha1};
use std::{collections::HashMap, sync::Arc};
use tokio::{
io::{AsyncBufReadExt, AsyncWriteExt, BufReader},
net::TcpStream,
sync::Mutex,
time::{Duration, timeout},
};
use super::WebSocket;
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);
// ---- 1. Read request line with timeout ----
let mut request_line = String::new();
timeout(Duration::from_secs(5), reader.read_line(&mut request_line)).await??;
let request_line = request_line.trim_end();
if !request_line.starts_with("GET") {
write_half
.write_all(
b"HTTP/1.1 405 Method Not Allowed\r\n\
Content-Length: 0\r\n\
Connection: close\r\n\r\n",
)
.await?;
write_half.shutdown().await?;
return Ok(());
}
// ---- 2. Read headers with timeout ----
let mut headers = HashMap::new();
loop {
let mut line = String::new();
timeout(Duration::from_secs(5), 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());
}
}
// ---- 3. Check if this is a WebSocket upgrade ----
let is_upgrade = headers
.get("upgrade")
.map(|v| v.eq_ignore_ascii_case("websocket"))
.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_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?;
write_half.write_all(body).await?;
write_half.flush().await?;
write_half.shutdown().await?;
return Ok(());
}
// ---- 4. Validate required headers ----
let key = headers.get("sec-websocket-key").ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidData, "Missing Sec-WebSocket-Key")
})?;
let version_ok = headers
.get("sec-websocket-version")
.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();
hasher.update(key.as_bytes());
hasher.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
let accept = Base64.encode(hasher.finalize());
// ---- 6. Send upgrade 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
);
write_half.write_all(response.as_bytes()).await?;
write_half.flush().await?;
Ok(())
}
impl WebSocket {
pub async fn handshake(mut stream: TcpStream) -> super::Result<Self> {
handle_websocket_handshake(&mut stream).await?;
let (read, write) = stream.into_split();
Ok(Self {
id: rand::random(),
reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)),
is_server: false,
})
}
/// Connect to a WebSocket server and perform the handshake
pub async fn connect(addr: &str, path: &str) -> super::Result<Self> {
// 1. TCP connect
let mut stream = TcpStream::connect(addr).await?;
// 2. Generate Sec-WebSocket-Key
let key_bytes: [u8; 16] = rand::random();
let key = base64::prelude::BASE64_STANDARD.encode(&key_bytes);
// 3. Send HTTP Upgrade request
let request = format!(
"GET {} HTTP/1.1\r\n\
Host: {}\r\n\
Upgrade: websocket\r\n\
Connection: Upgrade\r\n\
Sec-WebSocket-Key: {}\r\n\
Sec-WebSocket-Version: 13\r\n\
\r\n",
path, addr, key
);
stream.write_all(request.as_bytes()).await?;
stream.flush().await?;
// 4. Read HTTP response
let mut reader = BufReader::new(&mut stream);
let mut status_line = String::new();
timeout(
tokio::time::Duration::from_secs(5),
reader.read_line(&mut status_line),
)
.await??;
if !status_line.starts_with("HTTP/1.1 101") {
return Err(super::Error::HandshakeFailed(format!(
"Expected 101 Switching Protocols, got: {}",
status_line.trim_end()
)));
}
// Read headers
let mut sec_accept = None;
loop {
let mut line = String::new();
reader.read_line(&mut line).await?;
let line = line.trim_end();
if line.is_empty() {
break; // end of headers
}
if let Some((k, v)) = line.split_once(':') {
if k.eq_ignore_ascii_case("sec-websocket-accept") {
sec_accept = Some(v.trim().to_string());
}
}
}
// 5. Verify Sec-WebSocket-Accept
let expected = {
let mut sha1 = Sha1::new();
sha1.update(key.as_bytes());
sha1.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
base64::prelude::BASE64_STANDARD.encode(sha1.finalize())
};
if sec_accept.as_deref() != Some(expected.as_str()) {
return Err(super::Error::HandshakeFailed(
"Sec-WebSocket-Accept mismatch".into(),
));
}
// 6. Upgrade succeeded, split stream
let (read, write) = stream.into_split();
Ok(Self {
id: rand::random(),
reader: Arc::new(Mutex::new(read)),
writer: Arc::new(Mutex::new(write)),
is_server: true,
})
}
}
-251
View File
@@ -1,251 +0,0 @@
pub mod error;
pub mod handshake;
pub use error::{Error, Result};
use std::{
hash::{Hash, Hasher},
sync::Arc,
};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
sync::Mutex,
};
#[derive(Debug, Clone)]
pub enum Frame {
Text(String),
Binary(Vec<u8>),
Ping,
Pong,
Close,
}
pub struct WebSocket {
pub(crate) reader: Arc<Mutex<tokio::net::tcp::OwnedReadHalf>>,
pub(crate) writer: Arc<Mutex<tokio::net::tcp::OwnedWriteHalf>>,
pub(crate) id: u64,
pub(crate) is_server: bool,
}
impl Clone for WebSocket {
fn clone(&self) -> Self {
WebSocket {
reader: self.reader.clone(),
writer: self.writer.clone(),
is_server: self.is_server.clone(),
id: self.id,
}
}
}
impl PartialEq for WebSocket {
fn eq(&self, other: &Self) -> bool {
self.id == other.id
}
}
impl Eq for WebSocket {}
impl Hash for WebSocket {
fn hash<H: Hasher>(&self, state: &mut H) {
self.id.hash(state);
}
}
impl WebSocket {
async fn send_frame(&self, opcode: u8, payload: &[u8]) -> Result<()> {
let mut writer = self.writer.lock().await;
let mut header = Vec::with_capacity(10);
let mask_bit = if self.is_server { 0x80 } else { 0x00 };
header.push(0x80 | opcode); // FIN + opcode
let len = payload.len();
if len < 126 {
header.push((len as u8) | mask_bit);
} else if len <= 0xFFFF {
header.push(126 | mask_bit);
header.extend_from_slice(&(len as u16).to_be_bytes());
} else {
header.push(127 | mask_bit);
header.extend_from_slice(&(len as u64).to_be_bytes());
}
if self.is_server {
// Generate 4-byte mask key
let mask_key: [u8; 4] = rand::random();
header.extend_from_slice(&mask_key);
// Mask the payload
let mut masked_payload = payload.to_vec();
for i in 0..masked_payload.len() {
masked_payload[i] ^= mask_key[i % 4];
}
writer.write_all(&header).await?;
writer.write_all(&masked_payload).await?;
} else {
writer.write_all(&header).await?;
writer.write_all(payload).await?;
}
writer.flush().await?;
Ok(())
}
}
impl WebSocket {
pub async fn send(&self, msg: &str) -> Result<()> {
self.send_frame(0x1, msg.as_bytes()).await
}
pub async fn send_text_payload(&self, payload: &[u8]) -> Result<()> {
self.send_frame(0x1, payload).await
}
pub async fn send_bin(&self, payload: &[u8]) -> Result<()> {
self.send_frame(0x2, payload).await
}
pub async fn send_ping(&self) -> Result<()> {
self.send_frame(0x9, &[]).await
}
pub async fn send_pong(&self) -> Result<()> {
self.send_frame(0xA, &[]).await
}
pub async fn close(&self) -> Result<()> {
self.send_frame(0x8, &[]).await
}
pub fn start_ping_loop(&self) {
let s = self.clone();
tokio::task::spawn(async move {
let mut interval = tokio::time::interval(std::time::Duration::from_secs(15));
loop {
interval.tick().await;
if s.send_ping().await.is_err() {
break;
}
}
});
}
}
impl WebSocket {
/// Read a full WebSocket frame (handling masking and control frames)
/// Returns (opcode, payload)
pub async fn read_frame(&self) -> Result<(bool, u8, Vec<u8>)> {
let mut reader = self.reader.lock().await;
// --- 1. Read first 2-byte header ---
let mut header = [0u8; 2];
reader.read_exact(&mut header).await?;
let fin = header[0] & 0x80 != 0;
let opcode = header[0] & 0x0F;
let masked = header[1] & 0x80 != 0;
let mut payload_len = (header[1] & 0x7F) as u64;
// --- 2. Read extended payload length if necessary ---
if payload_len == 126 {
let mut buf = [0u8; 2];
reader.read_exact(&mut buf).await?;
payload_len = u16::from_be_bytes(buf) as u64;
} else if payload_len == 127 {
let mut buf = [0u8; 8];
reader.read_exact(&mut buf).await?;
payload_len = u64::from_be_bytes(buf);
}
let payload = if masked {
// --- 3. Read mask key ---
let mut mask = [0u8; 4];
reader.read_exact(&mut mask).await?;
let mut payload = vec![0u8; payload_len as usize];
if payload_len > 0 {
reader.read_exact(&mut payload).await?;
for i in 0..payload.len() {
payload[i] ^= mask[i % 4];
}
}
payload
} else {
// Per spec, client-to-server frames MUST be masked
if !self.is_server {
self.close().await.ok();
return Err(Error::InvalidFrame(
"Received unmasked frame from client".into(),
));
}
let mut payload = vec![0u8; payload_len as usize];
if payload_len > 0 {
reader.read_exact(&mut payload).await?;
}
payload
};
// --- 6. Return opcode + payload ---
Ok((fin, opcode, payload))
}
pub async fn read(&self) -> Result<Frame> {
let (fin, opcode, mut payload) = self.read_frame().await?;
if !fin {
// Continuation loop
while let (fin, o, mut p) = self.read_frame().await?
&& !fin
{
match o {
// Continuation
0x0 => payload.append(&mut p),
// Close
0x8 => {
self.close().await.ok();
}
// Ping
0x9 => {
self.send_pong().await.ok();
}
// Pong
0xA => {}
_ => {
self.close().await.ok();
return Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}")));
}
}
}
}
match opcode {
// Close
0x8 => {
self.close().await.ok();
Ok(Frame::Close)
}
// Ping
0x9 => {
self.send_pong().await.ok();
Ok(Frame::Ping)
}
// Pong
0xA => Ok(Frame::Pong),
// Text
0x1 => Ok(Frame::Text(String::from_utf8(payload)?)),
// Binary
0x2 => Ok(Frame::Binary(payload)),
_ => {
self.close().await.ok();
Err(Error::InvalidFrame(format!("Unknown opcode: {opcode}")))
}
}
}
}
+452
View File
@@ -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}");
}