0.2.0: transport-agnostic protocol on tokio-tungstenite, plus axum adapter #1
Generated
+958
-30
File diff suppressed because it is too large
Load Diff
+31
-14
@@ -1,21 +1,38 @@
|
||||
[package]
|
||||
name = "session-rs"
|
||||
version = "0.1.3"
|
||||
version = "0.2.0"
|
||||
edition = "2024"
|
||||
description = "A lightweight async WebSocket protocol"
|
||||
description = "A lightweight async request/response and notification protocol over WebSockets"
|
||||
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]
|
||||
base64 = "0.22.1"
|
||||
rand = "0.10.0"
|
||||
serde = { version = "1.0.228", features = ["serde_derive"] }
|
||||
futures-util = { version = "0.3.34", default-features = false, features = ["sink", "std"] }
|
||||
serde = { version = "1.0.228", features = ["derive"] }
|
||||
serde_json = "1.0.149"
|
||||
sha1 = "0.10.6"
|
||||
tokio = { version = "1.49.0", features = [
|
||||
"io-util",
|
||||
"macros",
|
||||
"net",
|
||||
"rt",
|
||||
"sync",
|
||||
"time",
|
||||
] }
|
||||
tokio = { version = "1.49.0", features = ["macros", "rt", "sync", "time"] }
|
||||
tokio-tungstenite = { version = "0.30.0", default-features = false, features = ["handshake"], optional = true }
|
||||
axum = { version = "0.8.9", default-features = false, features = ["ws"], optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
tokio = { version = "1.49.0", features = ["macros", "rt-multi-thread", "net", "time", "io-util"] }
|
||||
axum = { version = "0.8.9", features = ["ws"] }
|
||||
|
||||
[[example]]
|
||||
name = "axum"
|
||||
required-features = ["axum"]
|
||||
|
||||
[package.metadata.docs.rs]
|
||||
all-features = true
|
||||
|
||||
@@ -6,20 +6,12 @@
|
||||
|
||||
## Introduction
|
||||
|
||||
This library provides **type-safe WebSocket communication** with a request-response and notification system built on top of a flexible protocol.
|
||||
It ensures compile-time guarantees for message structure, reduces runtime errors, and simplifies building Rust client/server applications.
|
||||
`session-rs` is a small request/response + notification protocol that runs over WebSockets, with typed methods on both ends.
|
||||
|
||||
- **Dynamic Methods**: Each message includes a method enum for type safety.
|
||||
- **Typed Requests & Responses**: Automatic serialization and deserialization.
|
||||
- **Optional Notifications**: Send asynchronous notifications across sessions.
|
||||
|
||||
## Features
|
||||
|
||||
- Fully typed WebSocket sessions
|
||||
- Type-safe request/response mechanism
|
||||
- Optional typed notifications (Todo)
|
||||
- Lightweight, minimal runtime overhead
|
||||
- Async-first with Tokio support
|
||||
- **Typed methods**: requests, responses and errors are (de)serialized for you.
|
||||
- **Both directions**: either peer can send requests and notifications.
|
||||
- **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).
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -27,9 +19,16 @@ It ensures compile-time guarantees for message structure, reduces runtime errors
|
||||
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
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
@@ -41,44 +40,74 @@ impl Method for Data {
|
||||
type Response = 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
|
||||
#[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?;
|
||||
|
||||
server
|
||||
.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(())
|
||||
}).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
|
||||
|
||||
Every message is a JSON text frame with a `type` tag.
|
||||
|
||||
#### 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
|
||||
{ "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
|
||||
|
||||
The response `id` **must** remain the same as the request.
|
||||
A response **must** carry the id of the request it answers.
|
||||
|
||||
```json
|
||||
{ "type": "response", "id": 1, "result": "Hello from server" }
|
||||
```
|
||||
|
||||
#### Notifications
|
||||
|
||||
A notification is a method that doesn't need validation or output, it simply notifies a peer for a specific information
|
||||
#### Error response
|
||||
|
||||
```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" }
|
||||
```
|
||||
|
||||
@@ -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
@@ -1,5 +1,5 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use session_rs::{Method, session::Session};
|
||||
use session_rs::{Method, Session};
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
struct Data;
|
||||
@@ -13,7 +13,7 @@ impl Method for Data {
|
||||
|
||||
#[tokio::main(flavor = "current_thread")]
|
||||
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();
|
||||
|
||||
@@ -31,8 +31,6 @@ async fn main() -> session_rs::Result<()> {
|
||||
.await?
|
||||
);
|
||||
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(1000)).await;
|
||||
|
||||
session.close().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+39
@@ -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)),
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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};
|
||||
|
||||
#[cfg(feature = "axum")]
|
||||
mod axum;
|
||||
#[cfg(feature = "client")]
|
||||
mod client;
|
||||
#[cfg(feature = "server")]
|
||||
pub mod server;
|
||||
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 BoxFuture<'a, T = Option<(bool, serde_json::Value)>> =
|
||||
Pin<Box<dyn Future<Output = T> + Send + 'a>>;
|
||||
pub type MethodHandler = Arc<dyn Fn(u32, serde_json::Value) -> BoxFuture<'static> + Send + Sync>;
|
||||
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
|
||||
pub type BoxError = Box<dyn std::error::Error + 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 {
|
||||
const NAME: &'static str;
|
||||
type Request: Serialize + for<'de> Deserialize<'de> + Send + Sync;
|
||||
@@ -18,6 +37,7 @@ pub trait Method {
|
||||
type Error: Serialize + for<'de> Deserialize<'de>;
|
||||
}
|
||||
|
||||
/// Untyped method used internally for raw JSON values.
|
||||
pub struct GenericMethod;
|
||||
|
||||
impl Method for GenericMethod {
|
||||
@@ -29,18 +49,30 @@ impl Method for GenericMethod {
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum Error {
|
||||
WebSocket(ws::Error),
|
||||
/// The underlying transport (WebSocket) failed.
|
||||
Transport(BoxError),
|
||||
Json(serde_json::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 {
|
||||
fn from(value: ws::Error) -> Self {
|
||||
Self::WebSocket(value)
|
||||
impl std::fmt::Display for Error {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
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 {
|
||||
fn from(value: std::io::Error) -> Self {
|
||||
Self::Io(value)
|
||||
@@ -53,8 +85,8 @@ impl From<serde_json::Error> for Error {
|
||||
}
|
||||
}
|
||||
|
||||
impl From<tokio::sync::broadcast::error::RecvError> for Error {
|
||||
fn from(value: tokio::sync::broadcast::error::RecvError) -> Self {
|
||||
Self::RecvError(value)
|
||||
impl From<tokio::time::error::Elapsed> for Error {
|
||||
fn from(_: tokio::time::error::Elapsed) -> Self {
|
||||
Self::Timeout
|
||||
}
|
||||
}
|
||||
|
||||
+95
-33
@@ -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 {
|
||||
listener: TcpListener,
|
||||
config: ServerConfig,
|
||||
}
|
||||
|
||||
impl SessionServer {
|
||||
pub async fn bind(addr: &str) -> crate::Result<Self> {
|
||||
Ok(Self {
|
||||
listener: TcpListener::bind(addr).await?,
|
||||
})
|
||||
pub async fn bind(addr: impl ToSocketAddrs) -> crate::Result<Self> {
|
||||
Ok(Self::from_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)> {
|
||||
let (stream, addr) = self.listener.accept().await?;
|
||||
|
||||
let ws = WebSocket::handshake(stream).await?;
|
||||
|
||||
Ok((Session::from_ws(ws), addr))
|
||||
Ok((handshake(stream, &self.config).await?, 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<()>
|
||||
where
|
||||
F: Fn(Session, SocketAddr) -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = crate::Result<()>> + Send + 'static,
|
||||
{
|
||||
let conn_handler = Arc::new(on_conn);
|
||||
let on_conn = Arc::new(on_conn);
|
||||
|
||||
loop {
|
||||
let (stream, addr) = self.listener.accept().await?;
|
||||
let conn_handler = conn_handler.clone();
|
||||
let (stream, addr) = match self.listener.accept().await {
|
||||
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 {
|
||||
match timeout(
|
||||
tokio::time::Duration::from_secs(5),
|
||||
WebSocket::handshake(stream),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(ws)) => {
|
||||
let session = Session::from_ws(ws);
|
||||
session.start_receiver();
|
||||
let session = match handshake(stream, &config).await {
|
||||
Ok(session) => session,
|
||||
Err(e) => {
|
||||
eprintln!("Handshake failed from {addr}: {e}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(e) = conn_handler(session, addr).await {
|
||||
eprintln!("Connection error: {:?}", e);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
eprintln!("Handshake failed from {}: {:?}", addr, e);
|
||||
}
|
||||
Err(_) => {
|
||||
eprintln!("Handshake failed from {}: Handshake Timeout", addr);
|
||||
}
|
||||
if let Err(e) = on_conn(session.clone(), addr).await {
|
||||
eprintln!("Connection error from {addr}: {e}");
|
||||
let _ = session.close().await;
|
||||
return;
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
|
||||
+385
-152
@@ -1,14 +1,25 @@
|
||||
use std::collections::HashMap;
|
||||
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 tokio::sync::Mutex;
|
||||
use tokio::sync::broadcast;
|
||||
use tokio::time::timeout;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::{mpsc, oneshot, watch};
|
||||
|
||||
use crate::BoxFuture;
|
||||
use crate::{GenericMethod, Method, MethodHandler, ws::WebSocket};
|
||||
use crate::transport::{self, BoxSink, BoxStream, Frame};
|
||||
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)]
|
||||
#[serde(rename_all = "lowercase", tag = "type")]
|
||||
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 ws: WebSocket,
|
||||
id: Arc<Mutex<u32>>,
|
||||
methods: Arc<Mutex<HashMap<String, MethodHandler>>>,
|
||||
on_close_fn:
|
||||
Arc<Mutex<Option<Box<dyn Fn() -> BoxFuture<'static, Result<(), String>> + Send + Sync>>>>,
|
||||
tx: broadcast::Sender<(u32, bool, serde_json::Value)>,
|
||||
pong_tx: broadcast::Sender<()>,
|
||||
inner: Arc<Inner>,
|
||||
}
|
||||
|
||||
struct Inner {
|
||||
id: u64,
|
||||
next_request_id: AtomicU32,
|
||||
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 {
|
||||
pub fn clone(&self) -> Self {
|
||||
Self {
|
||||
ws: self.ws.clone(),
|
||||
id: self.id.clone(),
|
||||
methods: self.methods.clone(),
|
||||
on_close_fn: self.on_close_fn.clone(),
|
||||
tx: self.tx.clone(),
|
||||
pong_tx: self.pong_tx.clone(),
|
||||
/// Build a session over any frame sink/stream pair.
|
||||
///
|
||||
/// The transport is expected to answer pings itself (tungstenite and axum
|
||||
/// both do). Nothing is read until [`Session::start_receiver`] is called,
|
||||
/// so handlers can be registered first.
|
||||
pub fn from_transport<Si, St, SiE, StE>(sink: Si, stream: St) -> Self
|
||||
where
|
||||
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
|
||||
}
|
||||
|
||||
/// A process-unique id for this connection.
|
||||
pub fn id(&self) -> u64 {
|
||||
self.inner.id
|
||||
}
|
||||
|
||||
pub fn is_closed(&self) -> bool {
|
||||
self.inner.closed.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
/// Resolves once the session is closed, from either side.
|
||||
pub async fn closed(&self) {
|
||||
wait_closed(&mut self.inner.closed_tx.subscribe()).await;
|
||||
}
|
||||
}
|
||||
|
||||
impl Session {
|
||||
pub fn from_ws(ws: WebSocket) -> Self {
|
||||
let (tx, _) = broadcast::channel(8192);
|
||||
let (pong_tx, _) = broadcast::channel(16);
|
||||
|
||||
Self {
|
||||
ws,
|
||||
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> {
|
||||
Ok(Self::from_ws(WebSocket::connect(addr, path).await?))
|
||||
}
|
||||
}
|
||||
|
||||
impl Session {
|
||||
/// Start reading from the peer. Calling it again has no effect.
|
||||
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();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let mut closed = s.inner.closed_tx.subscribe();
|
||||
let mut pongs = s.inner.pongs.subscribe();
|
||||
|
||||
loop {
|
||||
match s.ws.read().await {
|
||||
Ok(crate::ws::Frame::Text(text)) => {
|
||||
let Ok(msg) = serde_json::from_str::<Message<GenericMethod>>(&text) else {
|
||||
continue;
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(interval) => {}
|
||||
_ = wait_closed(&mut closed) => return,
|
||||
}
|
||||
|
||||
pongs.mark_unchanged();
|
||||
|
||||
if s.inner.outgoing.send(Frame::Ping(Vec::new())).await.is_err() {
|
||||
break;
|
||||
}
|
||||
|
||||
match tokio::time::timeout(timeout, pongs.changed()).await {
|
||||
Ok(Ok(())) => {}
|
||||
_ => break,
|
||||
}
|
||||
}
|
||||
|
||||
s.shutdown().await;
|
||||
});
|
||||
}
|
||||
|
||||
async fn read_loop(self, mut stream: BoxStream, dispatch: mpsc::Sender<Incoming>) {
|
||||
let mut closed = self.inner.closed_tx.subscribe();
|
||||
|
||||
loop {
|
||||
let frame = tokio::select! {
|
||||
frame = stream.next() => frame,
|
||||
_ = wait_closed(&mut closed) => break,
|
||||
};
|
||||
|
||||
match msg {
|
||||
Message::Request { id, method, data } => {
|
||||
let handler = {
|
||||
let methods = s.methods.lock().await;
|
||||
methods.get(&method).cloned()
|
||||
match frame {
|
||||
Some(Ok(Frame::Text(text))) => {
|
||||
if !self.handle_text(&text, &dispatch).await {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Some(Ok(Frame::Pong(_))) => {
|
||||
self.inner.pongs.send_modify(|n| *n = n.wrapping_add(1));
|
||||
}
|
||||
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;
|
||||
};
|
||||
|
||||
if let Some(m) = handler {
|
||||
if let Some((err, res)) = (m)(id, data).await {
|
||||
if err {
|
||||
s.respond_error(id, res)
|
||||
.await
|
||||
.expect("Failed to respond");
|
||||
} else {
|
||||
s.respond(id, res).await.expect("Failed to respond");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
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 } => {
|
||||
s.tx.send((id, false, result)).unwrap();
|
||||
self.complete(id, Ok(result));
|
||||
return true;
|
||||
}
|
||||
Message::ErrorResponse { id, error } => {
|
||||
s.tx.send((id, true, error)).unwrap();
|
||||
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);
|
||||
}
|
||||
}
|
||||
Ok(crate::ws::Frame::Pong) => {
|
||||
let _ = s.pong_tx.send(());
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(_) => {
|
||||
s.trigger_close().await;
|
||||
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
pub fn start_ping(&self, interval: tokio::time::Duration, timeout_dur: tokio::time::Duration) {
|
||||
let s = self.clone();
|
||||
Incoming::Notification { method, data } => {
|
||||
let handler = self.inner.notifications.lock().unwrap().get(&method).cloned();
|
||||
|
||||
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;
|
||||
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<
|
||||
M: Method,
|
||||
Fut: Future<Output = Result<M::Response, M::Error>> + Send + 'static,
|
||||
@@ -158,57 +289,89 @@ impl Session {
|
||||
&self,
|
||||
handler: impl Fn(u32, M::Request) -> Fut + Send + Sync + 'static,
|
||||
) {
|
||||
let handler = Arc::new(handler);
|
||||
|
||||
self.methods.lock().await.insert(
|
||||
M::NAME.to_string(),
|
||||
Arc::new(move |id, value| {
|
||||
let handler = Arc::clone(&handler);
|
||||
let handler: RequestHandler = Arc::new(move |id, value| {
|
||||
let fut = serde_json::from_value::<M::Request>(value).map(|req| handler(id, req));
|
||||
|
||||
Box::pin(async move {
|
||||
Some(
|
||||
match handler(id, serde_json::from_value(value).ok()?).await {
|
||||
Ok(v) => (false, serde_json::to_value(v).ok()?),
|
||||
Err(v) => (true, serde_json::to_value(v).ok()?),
|
||||
match fut {
|
||||
Err(e) => Err(Value::from(format!("Invalid request data: {e}"))),
|
||||
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}"))
|
||||
})),
|
||||
},
|
||||
)
|
||||
}
|
||||
})
|
||||
}),
|
||||
);
|
||||
});
|
||||
|
||||
self.inner
|
||||
.requests
|
||||
.lock()
|
||||
.unwrap()
|
||||
.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)
|
||||
where
|
||||
Fut: Future<Output = Result<(), String>> + Send + 'static,
|
||||
{
|
||||
let handler = Arc::new(handler);
|
||||
|
||||
*self.on_close_fn.lock().await = Some(Box::new(move || {
|
||||
let handler = handler.clone();
|
||||
Box::pin(async move { handler().await })
|
||||
}));
|
||||
*self.inner.on_close.lock().unwrap() = Some(Arc::new(move || Box::pin(handler())));
|
||||
}
|
||||
}
|
||||
|
||||
impl Session {
|
||||
pub async fn send<M: Method>(&self, data: &Message<M>) -> crate::Result<()> {
|
||||
self.ws
|
||||
.send_text_payload(&serde_json::to_vec(&data)?)
|
||||
.await?;
|
||||
Ok(())
|
||||
pub async fn send<M: Method>(&self, msg: &Message<M>) -> crate::Result<()> {
|
||||
if self.is_closed() {
|
||||
return Err(Error::ConnectionClosed);
|
||||
}
|
||||
|
||||
pub async fn use_id(&self) -> u32 {
|
||||
let mut id = self.id.lock().await;
|
||||
*id += 1;
|
||||
*id
|
||||
let text = serde_json::to_string(msg)?;
|
||||
|
||||
self.inner
|
||||
.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>(
|
||||
&self,
|
||||
req: M::Request,
|
||||
) -> 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 {
|
||||
id,
|
||||
@@ -217,30 +380,27 @@ impl Session {
|
||||
})
|
||||
.await?;
|
||||
|
||||
let mut rx = self.tx.subscribe();
|
||||
|
||||
loop {
|
||||
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,
|
||||
Ok(match rx.await.map_err(|_| Error::ConnectionClosed)? {
|
||||
Ok(v) => Ok(serde_json::from_value(v)?),
|
||||
Err(e) => Err(serde_json::from_value(e)?),
|
||||
})
|
||||
}
|
||||
|
||||
/// [`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: serde_json::Value) -> crate::Result<()> {
|
||||
pub async fn respond_error(&self, to: u32, val: Value) -> crate::Result<()> {
|
||||
self.send::<GenericMethod>(&Message::ErrorResponse { id: to, error: val })
|
||||
.await
|
||||
}
|
||||
@@ -253,28 +413,101 @@ impl Session {
|
||||
.await
|
||||
}
|
||||
|
||||
async fn trigger_close(&self) {
|
||||
if let Some(handler) = self.on_close_fn.lock().await.as_ref() {
|
||||
/// Close the connection. Queued messages are flushed first.
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn close(&self) -> crate::Result<()> {
|
||||
let res = self.ws.close().await;
|
||||
self.trigger_close().await;
|
||||
Ok(res?)
|
||||
/// Removes a pending request if its `request` future is dropped early.
|
||||
struct PendingGuard<'a> {
|
||||
inner: &'a Inner,
|
||||
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 {
|
||||
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
|
||||
self.ws.id.hash(state);
|
||||
self.inner.id.hash(state);
|
||||
}
|
||||
}
|
||||
|
||||
impl PartialEq for Session {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.ws.id == other.ws.id
|
||||
self.inner.id == other.inner.id
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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))),
|
||||
)
|
||||
}
|
||||
@@ -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)),
|
||||
}
|
||||
}),
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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}")))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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