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();
}
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()))
})
}
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
}
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());
let res = match handshake_response(req, subprotocol) {
Ok(res) => res,
Err(err) => {
reject(req, err).await;
return Err(err.into());
}
};
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());
let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
unsafe {
io.set_config(CFG.with(Clone::clone));
}
io.stop_timer();
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)
}
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;
}
}
}
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
}
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(_)))));
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(_))));
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());
}
}