diff --git a/crates/engineioxide/src/service/mod.rs b/crates/engineioxide/src/service/mod.rs index 6b9d1c13..e5f1adf0 100644 --- a/crates/engineioxide/src/service/mod.rs +++ b/crates/engineioxide/src/service/mod.rs @@ -44,6 +44,7 @@ use tower_service::Service as TowerSvc; use crate::{ body::ResponseBody, config::EngineIoConfig, engine::EngineIo, handler::EngineIoHandler, + service::parser::RequestInfo, }; mod futures; @@ -136,6 +137,24 @@ where } } +impl TowerSvc<(St, http::Request<()>)> for EngineIoService +where + H: EngineIoHandler, + St: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + type Response = (); + type Error = parser::ParseError; + type Future = std::future::Ready>; + + fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn call(&mut self, (conn, req): (St, http::Request<()>)) -> Self::Future { + std::future::ready(ws_conn_inner(self.engine.clone(), conn, req)) + } +} + /// Hyper 1.0 Service implementation. impl HyperSvc> for EngineIoService where @@ -160,6 +179,38 @@ where } } +impl HyperSvc<(St, http::Request<()>)> for EngineIoService +where + H: EngineIoHandler, + St: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + type Response = (); + type Error = parser::ParseError; + type Future = std::future::Ready>; + + fn call(&self, (conn, req): (St, http::Request<()>)) -> Self::Future { + std::future::ready(ws_conn_inner(self.engine.clone(), conn, req)) + } +} + +fn ws_conn_inner( + engine: Arc>, + conn: St, + req: http::Request<()>, +) -> Result<(), parser::ParseError> +where + H: EngineIoHandler, + St: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + let RequestInfo { protocol, sid, .. } = RequestInfo::parse(&req, &engine.config)?; + + let (parts, _) = req.into_parts(); + let fut = crate::transport::ws::on_init(engine, conn, protocol, sid, parts); + tokio::spawn(fut); + + Ok(()) +} + #[cfg(feature = "__test_harness")] #[doc(hidden)] impl EngineIoService diff --git a/crates/engineioxide/src/service/parser.rs b/crates/engineioxide/src/service/parser.rs index 2cb5623a..2530dc27 100644 --- a/crates/engineioxide/src/service/parser.rs +++ b/crates/engineioxide/src/service/parser.rs @@ -127,7 +127,7 @@ pub struct RequestInfo { impl RequestInfo { /// Parse the request URI to extract the [`TransportType`](crate::service::TransportType) and the socket id. - fn parse(req: &Request, config: &EngineIoConfig) -> Result { + pub(crate) fn parse(req: &Request, config: &EngineIoConfig) -> Result { use ParseError::*; let query = req.uri().query().ok_or(UnknownTransport)?; diff --git a/crates/socketioxide/src/service.rs b/crates/socketioxide/src/service.rs index 55ea5225..131ee613 100644 --- a/crates/socketioxide/src/service.rs +++ b/crates/socketioxide/src/service.rs @@ -54,13 +54,34 @@ where type Error = , S> as TowerSvc>>::Error; type Future = , S> as TowerSvc>>::Future; - #[inline(always)] fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - self.engine_svc.poll_ready(cx) + TowerSvc::>::poll_ready(&mut self.engine_svc, cx) } - #[inline(always)] + fn call(&mut self, req: Request) -> Self::Future { - self.engine_svc.call(req) + TowerSvc::>::call(&mut self.engine_svc, req) + } +} + +type WsReq = (St, http::Request<()>); + +/// Service to create a websocket connection +impl TowerSvc<(St, http::Request<()>)> for SocketIoService +where + A: Adapter, + S: Clone, + St: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + type Response = , S> as HyperSvc>>::Response; + type Error = , S> as HyperSvc>>::Error; + type Future = , S> as HyperSvc>>::Future; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + TowerSvc::>::poll_ready(&mut self.engine_svc, cx) + } + + fn call(&mut self, req: WsReq) -> Self::Future { + TowerSvc::>::call(&mut self.engine_svc, req) } } @@ -80,7 +101,23 @@ where #[inline(always)] fn call(&self, req: Request) -> Self::Future { - self.engine_svc.call(req) + HyperSvc::>::call(&self.engine_svc, req) + } +} + +/// Service to create a websocket connection +impl HyperSvc> for SocketIoService +where + S: Clone, + A: Adapter, + St: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + type Response = , S> as HyperSvc>>::Response; + type Error = , S> as HyperSvc>>::Error; + type Future = , S> as HyperSvc>>::Future; + + fn call(&self, req: WsReq) -> Self::Future { + HyperSvc::>::call(&self.engine_svc, req) } }