ntex 4.0.0-beta.16

Framework for composable network services
//! `WebSockets` protocol support
use std::fmt;

pub use crate::ws::{CloseCode, CloseReason, Frame, Message, WsSink};

use crate::http::{ConnectionType, body::BodySize, error::ResponseError, h1, header};
use crate::io::{DispatchItem, IoConfig, Reason};
use crate::service::{Ctx, IntoService, Pipeline, Service, apply_fn};
use crate::web::HttpRequest;
use crate::ws::{self, error::HandshakeError, error::WsError, handshake};
use crate::{SharedCfg, rt, time::Seconds};

thread_local! {
    static CFG: SharedCfg = SharedCfg::new("WS")
        .add(IoConfig::new().set_keepalive_timeout(Seconds::ZERO))
        .into();
}

/// Returns an iterator over the subprotocols requested by the client
/// in the `Sec-Websocket-Protocol` header.
///
/// # Example
///
/// ```rust
/// use ntex::web::{self, HttpRequest, ws};
///
/// async fn service(frame: ws::Frame) -> Result<Option<ws::Message>, std::io::Error> {
///     // handle incoming frames
///     Ok(None)
/// }
///
/// async fn handler(req: HttpRequest) {
///     let chosen = ws::subprotocols(&req)
///         .find(|p| *p == "my-subprotocol");
///
///     if let Err(err) = ws::start(&req, chosen, service).await {
///         eprintln!("WebSocket error: {err:?}");
///     }
/// }
///
/// let app = web::App::default().route("/ws", web::get().to(handler));
/// ```
pub fn subprotocols(req: &HttpRequest) -> impl Iterator<Item = &str> {
    req.headers()
        .get_all(header::SEC_WEBSOCKET_PROTOCOL)
        .flat_map(|val| {
            val.to_str()
                .ok()
                .into_iter()
                .flat_map(|s| s.split(',').map(str::trim).filter(|s| !s.is_empty()))
        })
}

/// Start websocket service handling Frame messages with automatic control/stop logic,
/// including the chosen subprotocol in the response.
///
/// If `subprotocol` is `Some`, the `Sec-Websocket-Protocol` header will be included
/// in the response with the chosen protocol. The protocol must be a valid HTTP
/// token offered by the client. If `None`, the header is omitted.
///
/// If the handshake fails for an upgrade request, the handshake error response
/// is sent and the connection is closed.
///
/// # Example
///
/// ```rust
/// use ntex::web::{self, HttpRequest, ws};
///
/// async fn service(frame: ws::Frame) -> Result<Option<ws::Message>, std::io::Error> {
///     // handle incoming frames
///     Ok(None)
/// }
///
/// async fn handler(req: HttpRequest) {
///     let chosen = ws::subprotocols(&req)
///         .find(|p| *p == "graphql-ws" || *p == "graphql-transport-ws");
///
///     if let Err(err) = ws::start(&req, chosen, service).await {
///         eprintln!("WebSocket error: {err:?}");
///     }
/// }
///
/// let app = web::App::default().route("/ws", web::get().to(handler));
/// ```
pub async fn start<S>(
    req: &HttpRequest,
    subprotocol: Option<&str>,
    f: impl IntoService<S, WsSink, Frame>,
) -> Result<(), WsError<S::Error>>
where
    S: Service<WsSink, Frame, Res = Option<Message>> + 'static,
    S::Error: fmt::Debug,
{
    start_with(
        req,
        subprotocol,
        DispatchService {
            svc: f.into_service(),
        },
    )
    .await
}

/// Start websocket service handling raw `DispatchItem` messages requiring manual control/stop logic,
/// including the chosen subprotocol in the response.
///
/// If `subprotocol` is `Some`, the `Sec-Websocket-Protocol` header will be included
/// in the response with the chosen protocol. The protocol must be a valid HTTP
/// token offered by the client. If `None`, the header is omitted.
///
/// If the handshake fails for an upgrade request, the handshake error response
/// is sent and the connection is closed.
pub async fn start_with<S, Err>(
    req: &HttpRequest,
    subprotocol: Option<&str>,
    f: impl IntoService<S, WsSink, DispatchItem<WsSink>>,
) -> Result<(), WsError<Err>>
where
    S: Service<WsSink, DispatchItem<WsSink>, Res = Option<Message>, Error = WsError<Err>> + 'static,
    S::Error: fmt::Debug,
    Err: 'static,
{
    log::trace!("Start ws handshake verification for {:?}", req.path());

    // ws handshake
    let res = match handshake_response(req, subprotocol) {
        Ok(res) => res,
        Err(err) => {
            reject(req, err).await;
            return Err(err.into());
        }
    };

    // extract io
    let item = req
        .head()
        .take_io()
        .ok_or(HandshakeError::NoWebsocketUpgrade)?;
    let io = item.0;
    let codec = item.1;

    io.encode(h1::Message::Item((res, BodySize::Empty)), &codec)
        .map_err(|_| HandshakeError::NoWebsocketUpgrade)?;
    log::trace!("Ws handshake verification completed for {:?}", req.path());

    // create sink, it is also the dispatcher's codec
    let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());

    // create ws service
    // SAFETY: the HTTP dispatcher has transferred ownership of `io` to this
    // upgrade path, and no borrowed reference from `io.cfg()` is retained.
    unsafe {
        io.set_config(CFG.with(Clone::clone));
    }

    // the h1 dispatcher may have started a headers-read timer on this IO;
    // cancel it so DSP_TIMEOUT doesn't fire on the new WS dispatcher
    io.stop_timer();

    // start websockets service dispatcher
    let timeout_sink = sink.clone();
    let service = apply_fn(f.into_service(), async move |req, svc| {
        let result = svc.call(req).await;
        if matches!(&result, Ok(Some(Message::Close(_)))) {
            timeout_sink.start_close_timeout();
        }
        result
    });
    let result = crate::io::Dispatcher::new(io, sink.clone(), Pipeline::new(sink, service)).await;
    log::trace!("Ws handler is terminated: {result:?}");

    result
}

fn handshake_response(
    req: &HttpRequest,
    subprotocol: Option<&str>,
) -> Result<crate::http::Response<()>, HandshakeError> {
    let mut res = handshake(req.head())?;
    if let Some(protocol) = subprotocol {
        if !ws::is_token(protocol) || !subprotocols(req).any(|offered| offered == protocol) {
            return Err(HandshakeError::BadWebsocketProtocol);
        }
        res.set_header(header::SEC_WEBSOCKET_PROTOCOL, protocol);
    }
    Ok(res.build().into_parts().0)
}

/// Sends the handshake error response and closes the connection.
///
/// The I/O stream of an upgrade request belongs to the handler, a response
/// returned by the handler is not sent.
async fn reject(req: &HttpRequest, err: HandshakeError) {
    if let Some((io, codec)) = req.head().take_io() {
        let mut res = err.error_response().into_parts().0;
        res.head_mut().set_connection_type(ConnectionType::Close);
        if io
            .encode(h1::Message::Item((res, BodySize::Empty)), &codec)
            .is_ok()
        {
            let _ = io.shutdown().await;
        }
    }
}

/// Just a wrapper over a service handling WebSocket messages and propagating shutdown
struct DispatchService<S> {
    svc: S,
}

impl<S, E> Service<WsSink, DispatchItem<WsSink>> for DispatchService<S>
where
    S: Service<WsSink, Frame, Res = Option<Message>, Error = E>,
    E: fmt::Debug,
{
    type Res = Option<Message>;
    type Error = WsError<E>;

    crate::forward_ready!(WsSink, svc, WsError::Service);
    crate::forward_shutdown!(WsSink, svc);

    async fn call(
        &self,
        req: DispatchItem<WsSink>,
        ctx: Ctx<'_, Self, WsSink>,
    ) -> Result<Self::Res, Self::Error> {
        match req {
            DispatchItem::Item(item) => {
                let s = if matches!(item, Frame::Close(_)) {
                    Some(ctx.st().clone())
                } else {
                    None
                };
                let result = ctx.call(&self.svc, item).await.map_err(WsError::Service);
                if let Some(s) = s {
                    rt::spawn(async move { s.io().close() });
                }
                result
            }
            // a clean disconnect is not an error
            DispatchItem::Control(_) | DispatchItem::Stop(Reason::Io(None)) => Ok(None),
            DispatchItem::Stop(Reason::Service) => {
                Ok(Some(Message::Close(Some(ws::CloseReason {
                    code: ws::CloseCode::Away,
                    description: None,
                }))))
            }
            DispatchItem::Stop(Reason::KeepAlive) => Err(WsError::KeepAlive),
            DispatchItem::Stop(Reason::ReadTimeout) => Err(WsError::ReadTimeout),
            DispatchItem::Stop(Reason::WriteTimeout) => Err(WsError::WriteTimeout),
            DispatchItem::Stop(Reason::Decoder(e)) => {
                let sink = ctx.st();
                if !sink.is_closed() {
                    let reason = ws::CloseReason::from(ws::CloseCode::Protocol);
                    let _ = sink.send(Message::Close(Some(reason))).await;
                }
                Err(WsError::Protocol(e))
            }
            DispatchItem::Stop(Reason::Encoder(e)) => Err(WsError::Protocol(e)),
            DispatchItem::Stop(Reason::Io(e)) => Err(WsError::Disconnected(e)),
        }
    }
}

#[cfg(test)]
mod tests {
    use std::io;

    use super::*;
    use crate::io::{Control, Io, testing::IoTest};
    use crate::service::fn_service;
    use crate::ws::error::ProtocolError;

    #[crate::rt_test]
    async fn dispatch_service() {
        let (client, server) = IoTest::create();
        client.remote_buffer_cap(1 << 20);
        let io = Io::from(server);
        let sink = WsSink::new(io.get_ref(), ws::Codec::new(), crate::Cfg::default());
        let svc = Pipeline::new(
            sink.clone(),
            DispatchService {
                svc: fn_service(async |frame: Frame| match frame {
                    Frame::Text(_) => Err(io::Error::other("text")),
                    _ => Ok(Some(Message::Pong("pong".into()))),
                }),
            },
        );

        let res = svc.call(DispatchItem::Item(Frame::Ping("p".into()))).await;
        assert!(matches!(res, Ok(Some(Message::Pong(_)))));
        let res = svc.call(DispatchItem::Item(Frame::Text("t".into()))).await;
        assert!(matches!(res, Err(WsError::Service(_))));
        let res = svc
            .call(DispatchItem::Control(Control::WBackPressureEnabled))
            .await;
        assert!(matches!(res, Ok(None)));
        let res = svc.call(DispatchItem::Stop(Reason::Io(None))).await;
        assert!(matches!(res, Ok(None)));
        let res = svc.call(DispatchItem::Stop(Reason::Service)).await;
        assert!(matches!(
            res,
            Ok(Some(Message::Close(Some(ws::CloseReason {
                code: ws::CloseCode::Away,
                ..
            }))))
        ));
        let res = svc.call(DispatchItem::Stop(Reason::KeepAlive)).await;
        assert!(matches!(res, Err(WsError::KeepAlive)));
        let res = svc.call(DispatchItem::Stop(Reason::ReadTimeout)).await;
        assert!(matches!(res, Err(WsError::ReadTimeout)));
        let res = svc.call(DispatchItem::Stop(Reason::WriteTimeout)).await;
        assert!(matches!(res, Err(WsError::WriteTimeout)));
        let res = svc
            .call(DispatchItem::Stop(Reason::Encoder(
                ProtocolError::UnmaskedFrame,
            )))
            .await;
        assert!(matches!(
            res,
            Err(WsError::Protocol(ProtocolError::UnmaskedFrame))
        ));
        let res = svc
            .call(DispatchItem::Stop(Reason::Io(Some(io::Error::other("io")))))
            .await;
        assert!(matches!(res, Err(WsError::Disconnected(Some(_)))));

        // decoder error sends close frame
        assert!(!sink.is_closed());
        let res = svc
            .call(DispatchItem::Stop(Reason::Decoder(
                ProtocolError::MaskedFrame,
            )))
            .await;
        assert!(matches!(
            res,
            Err(WsError::Protocol(ProtocolError::MaskedFrame))
        ));
        assert!(sink.is_closed());
        let res = svc
            .call(DispatchItem::Stop(Reason::Decoder(
                ProtocolError::MaskedFrame,
            )))
            .await;
        assert!(matches!(res, Err(WsError::Protocol(_))));

        // close frame closes io
        let res = svc.call(DispatchItem::Item(Frame::Close(None))).await;
        assert!(matches!(res, Ok(Some(Message::Pong(_)))));
        io.on_disconnect().await;
        assert!(io.is_closed());
    }
}