diff --git a/Cargo.toml b/Cargo.toml
index b908b562..98561cf3 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -12,6 +12,9 @@ readme = "README.md"
repository = "https://github.com/davidpdrsn/tower-web"
version = "0.1.0"
+[features]
+ws = ["tokio-tungstenite", "sha-1", "base64"]
+
[dependencies]
async-trait = "0.1"
bytes = "1.0"
@@ -28,16 +31,32 @@ tokio = { version = "1", features = ["time"] }
tower = { version = "0.4", features = ["util", "buffer"] }
tower-http = { version = "0.1", features = ["add-extension"] }
+# optional dependencies
+tokio-tungstenite = { optional = true, version = "0.14" }
+sha-1 = { optional = true, version = "0.9.6" }
+base64 = { optional = true, version = "0.13" }
+
[dev-dependencies]
hyper = { version = "0.14", features = ["full"] }
reqwest = { version = "0.11", features = ["json", "stream"] }
serde = { version = "1.0", features = ["derive"] }
tokio = { version = "1.6.1", features = ["macros", "rt", "rt-multi-thread"] }
-tower = { version = "0.4", features = ["util", "make", "timeout", "limit", "load-shed", "steer"] }
tracing = "0.1"
tracing-subscriber = "0.2"
uuid = "0.8"
+[dev-dependencies.tower]
+version = "0.4"
+features = [
+ "util",
+ "make",
+ "timeout",
+ "limit",
+ "load-shed",
+ "steer",
+ "filter",
+]
+
[dev-dependencies.tower-http]
version = "0.1"
features = [
diff --git a/examples/websocket.rs b/examples/websocket.rs
new file mode 100644
index 00000000..f58ce0ad
--- /dev/null
+++ b/examples/websocket.rs
@@ -0,0 +1,63 @@
+//! Example websocket server.
+//!
+//! Run with
+//!
+//! ```
+//! RUST_LOG=tower_http=debug,key_value_store=trace \
+//! cargo run \
+//! --features ws \
+//! --example websocket
+//! ```
+
+use http::StatusCode;
+use hyper::Server;
+use std::net::SocketAddr;
+use tower::make::Shared;
+use tower_http::{
+ services::ServeDir,
+ trace::{DefaultMakeSpan, TraceLayer},
+};
+use tower_web::{
+ prelude::*,
+ routing::nest,
+ service::ServiceExt,
+ ws::{ws, Message, WebSocket},
+};
+
+#[tokio::main]
+async fn main() {
+ tracing_subscriber::fmt::init();
+
+ // build our application with some routes
+ let app = nest(
+ "/",
+ ServeDir::new("examples/websocket")
+ .append_index_html_on_directories(true)
+ .handle_error(|error| (StatusCode::INTERNAL_SERVER_ERROR, error.to_string())),
+ )
+ // 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))
+ // logging so we can see whats going on
+ .layer(
+ TraceLayer::new_for_http().make_span_with(DefaultMakeSpan::default().include_headers(true)),
+ );
+
+ // run it with hyper
+ let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
+ tracing::debug!("listening on {}", addr);
+ let server = Server::bind(&addr).serve(Shared::new(app));
+ server.await.unwrap();
+}
+
+async fn handle_socket(mut socket: WebSocket) {
+ if let Some(msg) = socket.recv().await {
+ let msg = msg.unwrap();
+ println!("Client says: {:?}", msg);
+ }
+
+ loop {
+ socket.send(Message::text("Hi!")).await.unwrap();
+ tokio::time::sleep(std::time::Duration::from_secs(3)).await;
+ }
+}
diff --git a/examples/websocket/index.html b/examples/websocket/index.html
new file mode 100644
index 00000000..390bb86b
--- /dev/null
+++ b/examples/websocket/index.html
@@ -0,0 +1 @@
+
diff --git a/examples/websocket/script.js b/examples/websocket/script.js
new file mode 100644
index 00000000..3f166736
--- /dev/null
+++ b/examples/websocket/script.js
@@ -0,0 +1,9 @@
+const socket = new WebSocket('ws://localhost:3000/ws');
+
+socket.addEventListener('open', function (event) {
+ socket.send('Hello Server!');
+});
+
+socket.addEventListener('message', function (event) {
+ console.log('Message from server ', event.data);
+});
diff --git a/src/lib.rs b/src/lib.rs
index ccf9c2fb..031f8e5c 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -610,6 +610,10 @@ pub mod response;
pub mod routing;
pub mod service;
+#[cfg(feature = "ws")]
+#[cfg_attr(docsrs, doc(cfg(feature = "ws")))]
+pub mod ws;
+
#[cfg(test)]
mod tests;
diff --git a/src/response.rs b/src/response.rs
index b32504f6..dea64edd 100644
--- a/src/response.rs
+++ b/src/response.rs
@@ -147,6 +147,17 @@ where
}
}
+impl IntoResponse for (HeaderMap, T)
+where
+ T: Into,
+{
+ fn into_response(self) -> Response {
+ let mut res = Response::new(self.1.into());
+ *res.headers_mut() = self.0;
+ res
+ }
+}
+
impl IntoResponse for (StatusCode, HeaderMap, T)
where
T: Into,
diff --git a/src/routing.rs b/src/routing.rs
index 9d4c3e45..19579e45 100644
--- a/src/routing.rs
+++ b/src/routing.rs
@@ -770,12 +770,16 @@ where
fn strip_prefix(uri: &Uri, prefix: &str) -> Uri {
let path_and_query = if let Some(path_and_query) = uri.path_and_query() {
- let new_path = if let Some(path) = path_and_query.path().strip_prefix(prefix) {
+ let mut new_path = if let Some(path) = path_and_query.path().strip_prefix(prefix) {
path
} else {
path_and_query.path()
};
+ if new_path.is_empty() {
+ new_path = "/";
+ }
+
if let Some(query) = path_and_query.query() {
Some(
format!("{}?{}", new_path, query)
diff --git a/src/ws/future.rs b/src/ws/future.rs
new file mode 100644
index 00000000..9dd56223
--- /dev/null
+++ b/src/ws/future.rs
@@ -0,0 +1,68 @@
+//! Future types.
+
+use bytes::Bytes;
+use http::{HeaderValue, Response, StatusCode};
+use http_body::Full;
+use sha1::{Digest, Sha1};
+use std::{
+ convert::Infallible,
+ future::Future,
+ pin::Pin,
+ task::{Context, Poll},
+};
+
+/// Response future for [`WebSocketUpgrade`](super::WebSocketUpgrade).
+#[derive(Debug)]
+pub struct ResponseFuture(Result