1use 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
19pub 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
54pub 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
104pub 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 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 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 let sink = WsSink::new(io.get_ref(), ws::Codec::new(), io.shared().get());
148
149 unsafe {
153 io.set_config(CFG.with(Clone::clone));
154 }
155
156 io.stop_timer();
159
160 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
189async 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
206struct 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 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 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 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}