Change WebSocket API to use an extractor (#121)

Fixes https://github.com/tokio-rs/axum/issues/111

Example usage:

```rust
use axum::{
    prelude::*,
    extract::ws::{WebSocketUpgrade, WebSocket},
    response::IntoResponse,
};

let app = route("/ws", get(handler));

async fn handler(ws: WebSocketUpgrade) -> impl IntoResponse {
    ws.on_upgrade(handle_socket)
}

async fn handle_socket(mut socket: WebSocket) {
    while let Some(msg) = socket.recv().await {
        let msg = if let Ok(msg) = msg {
            msg
        } else {
            // client disconnected
            return;
        };

        if socket.send(msg).await.is_err() {
            // client disconnected
            return;
        }
    }
}
```
This commit is contained in:
David Pedersen
2021-08-07 17:26:23 +02:00
committed by GitHub
parent 404a3b5e8a
commit 4194cf70da
8 changed files with 619 additions and 652 deletions
+10 -6
View File
@@ -14,9 +14,9 @@ use futures::{sink::SinkExt, stream::StreamExt};
use tokio::sync::broadcast;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::prelude::*;
use axum::response::Html;
use axum::ws::{ws, Message, WebSocket};
use axum::response::{Html, IntoResponse};
use axum::AddExtensionLayer;
// Our shared state
@@ -33,7 +33,7 @@ async fn main() {
let app_state = Arc::new(AppState { user_set, tx });
let app = route("/", get(index))
.route("/websocket", ws(websocket))
.route("/websocket", get(websocket_handler))
.layer(AddExtensionLayer::new(app_state));
let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
@@ -44,10 +44,14 @@ async fn main() {
.unwrap();
}
async fn websocket(
stream: WebSocket,
async fn websocket_handler(
ws: WebSocketUpgrade,
extract::Extension(state): extract::Extension<Arc<AppState>>,
) {
) -> impl IntoResponse {
ws.on_upgrade(|socket| websocket(socket, state))
}
async fn websocket(stream: WebSocket, state: Arc<AppState>) {
// By splitting we can send and receive at the same time.
let (mut sender, mut receiver) = stream.split();
+13 -7
View File
@@ -7,11 +7,14 @@
//! ```
use axum::{
extract::TypedHeader,
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
TypedHeader,
},
prelude::*,
response::IntoResponse,
routing::nest,
service::ServiceExt,
ws::{ws, Message, WebSocket},
};
use http::StatusCode;
use std::net::SocketAddr;
@@ -44,7 +47,7 @@ async fn main() {
)
// routes are matched from bottom to top, so we have to put `nest` at the
// top since it matches all routes
.route("/ws", ws(handle_socket))
.route("/ws", get(ws_handler))
// logging so we can see whats going on
.layer(
TraceLayer::new_for_http().make_span_with(DefaultMakeSpan::default().include_headers(true)),
@@ -59,15 +62,18 @@ async fn main() {
.unwrap();
}
async fn handle_socket(
mut socket: WebSocket,
// websocket handlers can also use extractors
async fn ws_handler(
ws: WebSocketUpgrade,
user_agent: Option<TypedHeader<headers::UserAgent>>,
) {
) -> impl IntoResponse {
if let Some(TypedHeader(user_agent)) = user_agent {
println!("`{}` connected", user_agent.as_str());
}
ws.on_upgrade(handle_socket)
}
async fn handle_socket(mut socket: WebSocket) {
if let Some(msg) = socket.recv().await {
if let Ok(msg) = msg {
println!("Client says: {:?}", msg);