Skip to main content

ntex/web/
ws.rs

1//! `WebSockets` protocol support
2use std::fmt;
3
4pub use crate::ws::{CloseCode, CloseReason, Frame, Message, WsSink};
5
6use crate::http::{ConnectionType, body::BodySize, error::ResponseError, h1, header};
7use crate::io::{DispatchItem, IoConfig, Reason};
8use crate::service::{Ctx, IntoService, Pipeline, Service, apply_fn};
9use crate::web::HttpRequest;
10use crate::ws::{self, error::HandshakeError, error::WsError, handshake};
11use crate::{SharedCfg, rt, time::Seconds};
12
13thread_local! {
14    static CFG: SharedCfg = SharedCfg::new("WS")
15        .add(IoConfig::new().set_keepalive_timeout(Seconds::ZERO))
16        .into();
17}
18
19/// Returns an iterator over the subprotocols requested by the client
20/// in the `Sec-Websocket-Protocol` header.
21///
22/// # Example
23///
24/// ```rust
25/// use ntex::web::{self, HttpRequest, ws};
26///
27/// async fn service(frame: ws::Frame) -> Result<Option<ws::Message>, std::io::Error> {
28///     // handle incoming frames
29///     Ok(None)
30/// }
31///
32/// async fn handler(req: HttpRequest) {
33///     let chosen = ws::subprotocols(&req)
34///         .find(|p| *p == "my-subprotocol");
35///
36///     if let Err(err) = ws::start(&req, chosen, service).await {
37///         eprintln!("WebSocket error: {err:?}");
38///     }
39/// }
40///
41/// let app = web::App::default().route("/ws", web::get().to(handler));
42/// ```
43pub fn subprotocols(req: &HttpRequest) -> impl Iterator<Item = &str> {
44    req.headers()
45        .get_all(header::SEC_WEBSOCKET_PROTOCOL)
46        .flat_map(|val| {
47            val.to_str()
48                .ok()
49                .into_iter()
50                .flat_map(|s| s.split(',').map(str::trim).filter(|s| !s.is_empty()))
51        })
52}
53
54/// Start websocket service handling Frame messages with automatic control/stop logic,
55/// including the chosen subprotocol in the response.
56///
57/// If `subprotocol` is `Some`, the `Sec-Websocket-Protocol` header will be included
58/// in the response with the chosen protocol. The protocol must be a valid HTTP
59/// token offered by the client. If `None`, the header is omitted.
60///
61/// If the handshake fails for an upgrade request, the handshake error response
62/// is sent and the connection is closed.
63///
64/// # Example
65///
66/// ```rust
67/// use ntex::web::{self, HttpRequest, ws};
68///
69/// async fn service(frame: ws::Frame) -> Result<Option<ws::Message>, std::io::Error> {
70///     // handle incoming frames
71///     Ok(None)
72/// }
73///
74/// async fn handler(req: HttpRequest) {
75///     let chosen = ws::subprotocols(&req)
76///         .find(|p| *p == "graphql-ws" || *p == "graphql-transport-ws");
77///
78///     if let Err(err) = ws::start(&req, chosen, service).await {
79///         eprintln!("WebSocket error: {err:?}");
80///     }
81/// }
82///
83/// let app = web::App::default().route("/ws", web::get().to(handler));
84/// ```
85pub async fn start<S>(
86    req: &HttpRequest,
87    subprotocol: Option<&str>,
88    f: impl IntoService<S, WsSink, Frame>,
89) -> Result<(), WsError<S::Error>>
90where
91    S: Service<WsSink, Frame, Res = Option<Message>> + 'static,
92    S::Error: fmt::Debug,
93{
94    start_with(
95        req,
96        subprotocol,
97        DispatchService {
98            svc: f.into_service(),
99        },
100    )
101    .await
102}
103
104/// Start websocket service handling raw `DispatchItem` messages requiring manual control/stop logic,
105/// including the chosen subprotocol in the response.
106///
107/// If `subprotocol` is `Some`, the `Sec-Websocket-Protocol` header will be included
108/// in the response with the chosen protocol. The protocol must be a valid HTTP
109/// token offered by the client. If `None`, the header is omitted.
110///
111/// If the handshake fails for an upgrade request, the handshake error response
112/// is sent and the connection is closed.
113pub async fn start_with<S, Err>(
114    req: &HttpRequest,
115    subprotocol: Option<&str>,
116    f: impl IntoService<S, WsSink, DispatchItem<WsSink>>,
117) -> Result<(), WsError<Err>>
118where
119    S: Service<WsSink, DispatchItem<WsSink>, Res = Option<Message>, Error = WsError<Err>> + 'static,
120    S::Error: fmt::Debug,
121    Err: 'static,
122{
123    log::trace!("Start ws handshake verification for {:?}", req.path());
124
125    // ws handshake
126    let res = match handshake_response(req, subprotocol) {
127        Ok(res) => res,
128        Err(err) => {
129            reject(req, err).await;
130            return Err(err.into());
131        }
132    };
133
134    // extract io
135    let item = req
136        .head()
137        .take_io()
138        .ok_or(HandshakeError::NoWebsocketUpgrade)?;
139    let io = item.0;
140    let codec = item.1;
141
142    io.encode(h1::Message::Item((res, BodySize::Empty)), &codec)
143        .map_err(|_| HandshakeError::NoWebsocketUpgrade)?;
144    log::trace!("Ws handshake verification completed for {:?}", req.path());
145
146    // create sink, it is also the dispatcher's codec
147    let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
148
149    // create ws service
150    // SAFETY: the HTTP dispatcher has transferred ownership of `io` to this
151    // upgrade path, and no borrowed reference from `io.cfg()` is retained.
152    unsafe {
153        io.set_config(CFG.with(Clone::clone));
154    }
155
156    // the h1 dispatcher may have started a headers-read timer on this IO;
157    // cancel it so DSP_TIMEOUT doesn't fire on the new WS dispatcher
158    io.stop_timer();
159
160    // start websockets service dispatcher
161    let timeout_sink = sink.clone();
162    let service = apply_fn(f.into_service(), async move |req, svc| {
163        let result = svc.call(req).await;
164        if matches!(&result, Ok(Some(Message::Close(_)))) {
165            timeout_sink.start_close_timeout();
166        }
167        result
168    });
169    let result = crate::io::Dispatcher::new(io, sink.clone(), Pipeline::new(sink, service)).await;
170    log::trace!("Ws handler is terminated: {result:?}");
171
172    result
173}
174
175fn handshake_response(
176    req: &HttpRequest,
177    subprotocol: Option<&str>,
178) -> Result<crate::http::Response<()>, HandshakeError> {
179    let mut res = handshake(req.head())?;
180    if let Some(protocol) = subprotocol {
181        if !ws::is_token(protocol) || !subprotocols(req).any(|offered| offered == protocol) {
182            return Err(HandshakeError::BadWebsocketProtocol);
183        }
184        res.set_header(header::SEC_WEBSOCKET_PROTOCOL, protocol);
185    }
186    Ok(res.build().into_parts().0)
187}
188
189/// Sends the handshake error response and closes the connection.
190///
191/// The I/O stream of an upgrade request belongs to the handler, a response
192/// returned by the handler is not sent.
193async fn reject(req: &HttpRequest, err: HandshakeError) {
194    if let Some((io, codec)) = req.head().take_io() {
195        let mut res = err.error_response().into_parts().0;
196        res.head_mut().set_connection_type(ConnectionType::Close);
197        if io
198            .encode(h1::Message::Item((res, BodySize::Empty)), &codec)
199            .is_ok()
200        {
201            let _ = io.shutdown().await;
202        }
203    }
204}
205
206/// Just a wrapper over a service handling WebSocket messages and propagating shutdown
207struct DispatchService<S> {
208    svc: S,
209}
210
211impl<S, E> Service<WsSink, DispatchItem<WsSink>> for DispatchService<S>
212where
213    S: Service<WsSink, Frame, Res = Option<Message>, Error = E>,
214    E: fmt::Debug,
215{
216    type Res = Option<Message>;
217    type Error = WsError<E>;
218
219    crate::forward_ready!(WsSink, svc, WsError::Service);
220    crate::forward_shutdown!(WsSink, svc);
221
222    async fn call(
223        &self,
224        req: DispatchItem<WsSink>,
225        ctx: Ctx<'_, Self, WsSink>,
226    ) -> Result<Self::Res, Self::Error> {
227        match req {
228            DispatchItem::Item(item) => {
229                let s = if matches!(item, Frame::Close(_)) {
230                    Some(ctx.st().clone())
231                } else {
232                    None
233                };
234                let result = ctx.call(&self.svc, item).await.map_err(WsError::Service);
235                if let Some(s) = s {
236                    rt::spawn(async move { s.io().close() });
237                }
238                result
239            }
240            // a clean disconnect is not an error
241            DispatchItem::Control(_) | DispatchItem::Stop(Reason::Io(None)) => Ok(None),
242            DispatchItem::Stop(Reason::Service) => {
243                Ok(Some(Message::Close(Some(ws::CloseReason {
244                    code: ws::CloseCode::Away,
245                    description: None,
246                }))))
247            }
248            DispatchItem::Stop(Reason::KeepAlive) => Err(WsError::KeepAlive),
249            DispatchItem::Stop(Reason::ReadTimeout) => Err(WsError::ReadTimeout),
250            DispatchItem::Stop(Reason::WriteTimeout) => Err(WsError::WriteTimeout),
251            DispatchItem::Stop(Reason::Decoder(e)) => {
252                let sink = ctx.st();
253                if !sink.is_closed() {
254                    let reason = ws::CloseReason::from(ws::CloseCode::Protocol);
255                    let _ = sink.send(Message::Close(Some(reason))).await;
256                }
257                Err(WsError::Protocol(e))
258            }
259            DispatchItem::Stop(Reason::Encoder(e)) => Err(WsError::Protocol(e)),
260            DispatchItem::Stop(Reason::Io(e)) => Err(WsError::Disconnected(e)),
261        }
262    }
263}
264
265#[cfg(test)]
266mod tests {
267    use std::io;
268
269    use super::*;
270    use crate::io::{Control, Io, testing::IoTest};
271    use crate::service::fn_service;
272    use crate::ws::error::ProtocolError;
273
274    #[crate::rt_test]
275    async fn dispatch_service() {
276        let (client, server) = IoTest::create();
277        client.remote_buffer_cap(1 << 20);
278        let io = Io::from(server);
279        let sink = WsSink::new(io.get_ref(), ws::Codec::new(), crate::Cfg::default());
280        let svc = Pipeline::new(
281            sink.clone(),
282            DispatchService {
283                svc: fn_service(async |frame: Frame| match frame {
284                    Frame::Text(_) => Err(io::Error::other("text")),
285                    _ => Ok(Some(Message::Pong("pong".into()))),
286                }),
287            },
288        );
289
290        let res = svc.call(DispatchItem::Item(Frame::Ping("p".into()))).await;
291        assert!(matches!(res, Ok(Some(Message::Pong(_)))));
292        let res = svc.call(DispatchItem::Item(Frame::Text("t".into()))).await;
293        assert!(matches!(res, Err(WsError::Service(_))));
294        let res = svc
295            .call(DispatchItem::Control(Control::WBackPressureEnabled))
296            .await;
297        assert!(matches!(res, Ok(None)));
298        let res = svc.call(DispatchItem::Stop(Reason::Io(None))).await;
299        assert!(matches!(res, Ok(None)));
300        let res = svc.call(DispatchItem::Stop(Reason::Service)).await;
301        assert!(matches!(
302            res,
303            Ok(Some(Message::Close(Some(ws::CloseReason {
304                code: ws::CloseCode::Away,
305                ..
306            }))))
307        ));
308        let res = svc.call(DispatchItem::Stop(Reason::KeepAlive)).await;
309        assert!(matches!(res, Err(WsError::KeepAlive)));
310        let res = svc.call(DispatchItem::Stop(Reason::ReadTimeout)).await;
311        assert!(matches!(res, Err(WsError::ReadTimeout)));
312        let res = svc.call(DispatchItem::Stop(Reason::WriteTimeout)).await;
313        assert!(matches!(res, Err(WsError::WriteTimeout)));
314        let res = svc
315            .call(DispatchItem::Stop(Reason::Encoder(
316                ProtocolError::UnmaskedFrame,
317            )))
318            .await;
319        assert!(matches!(
320            res,
321            Err(WsError::Protocol(ProtocolError::UnmaskedFrame))
322        ));
323        let res = svc
324            .call(DispatchItem::Stop(Reason::Io(Some(io::Error::other("io")))))
325            .await;
326        assert!(matches!(res, Err(WsError::Disconnected(Some(_)))));
327
328        // decoder error sends close frame
329        assert!(!sink.is_closed());
330        let res = svc
331            .call(DispatchItem::Stop(Reason::Decoder(
332                ProtocolError::MaskedFrame,
333            )))
334            .await;
335        assert!(matches!(
336            res,
337            Err(WsError::Protocol(ProtocolError::MaskedFrame))
338        ));
339        assert!(sink.is_closed());
340        let res = svc
341            .call(DispatchItem::Stop(Reason::Decoder(
342                ProtocolError::MaskedFrame,
343            )))
344            .await;
345        assert!(matches!(res, Err(WsError::Protocol(_))));
346
347        // close frame closes io
348        let res = svc.call(DispatchItem::Item(Frame::Close(None))).await;
349        assert!(matches!(res, Ok(Some(Message::Pong(_)))));
350        io.on_disconnect().await;
351        assert!(io.is_closed());
352    }
353}