Skip to main content

serverkit_hyper/
lib.rs

1#![forbid(unsafe_code)]
2
3use std::{
4    convert::Infallible,
5    future::Future,
6    io,
7    net::{TcpListener, ToSocketAddrs},
8    pin::Pin,
9    rc::Rc,
10    task::{Context, Poll},
11};
12
13use hyper::{
14    Request as HyperRequest, Response as HyperResponse,
15    body::{Body as HyperBody, Bytes, Frame, Incoming, SizeHint},
16    header::{CONTENT_LENGTH, HeaderName, HeaderValue},
17    rt::Executor,
18    service::service_fn,
19};
20use hyper_util::{rt::TokioIo, server::conn::auto};
21
22#[cfg(feature = "websocket")]
23use futures_util::{Sink, Stream};
24#[cfg(feature = "websocket")]
25use tokio_tungstenite::{
26    WebSocketStream,
27    tungstenite::{
28        Message,
29        handshake::derive_accept_key,
30        protocol::{CloseFrame, Role, frame::coding::CloseCode},
31    },
32};
33
34use serverkit::{
35    Chunk, Headers, Listener, Method, Request, RequestStream, Response, ResponseBody, Router,
36    StreamError,
37};
38
39#[cfg(feature = "websocket")]
40use serverkit::{
41    WebSocket, WebSocketError, WebSocketMessage,
42    adapter::{WebSocketIo, WebSocketPlan},
43};
44
45pub struct Http {
46    listener: TcpListener,
47}
48
49impl Http {
50    pub fn bind(address: impl ToSocketAddrs) -> io::Result<Self> {
51        TcpListener::bind(address).map(Self::from_listener)
52    }
53
54    pub fn from_listener(listener: TcpListener) -> Self {
55        Self { listener }
56    }
57}
58
59impl Listener for Http {
60    type Output = io::Result<()>;
61
62    fn serve(self, router: Router) -> Self::Output {
63        serve(router, self.listener)
64    }
65}
66
67fn serve(router: Router, listener: TcpListener) -> io::Result<()> {
68    listener.set_nonblocking(true)?;
69
70    let runtime = tokio::runtime::Builder::new_current_thread()
71        .enable_io()
72        .build()?;
73    let tasks = tokio::task::LocalSet::new();
74
75    tasks.block_on(&runtime, serve_connections(router, listener))
76}
77
78async fn serve_connections(router: Router, listener: TcpListener) -> io::Result<()> {
79    let listener = tokio::net::TcpListener::from_std(listener)?;
80    let router = Rc::new(router);
81
82    loop {
83        let (connection, address) = listener.accept().await?;
84        let router = Rc::clone(&router);
85
86        tokio::task::spawn_local(async move {
87            serve_connection(router, connection, address).await;
88        });
89    }
90}
91
92async fn serve_connection(
93    router: Rc<Router>,
94    connection: tokio::net::TcpStream,
95    address: std::net::SocketAddr,
96) {
97    let service = service_fn(move |request| {
98        let router = Rc::clone(&router);
99
100        async move { Ok::<_, Infallible>(handle_request(router, request, address).await) }
101    });
102
103    let builder = auto::Builder::new(LocalExecutor);
104
105    #[cfg(feature = "websocket")]
106    let _result = builder
107        .serve_connection_with_upgrades(TokioIo::new(connection), service)
108        .await;
109
110    #[cfg(not(feature = "websocket"))]
111    let _result = builder
112        .serve_connection(TokioIo::new(connection), service)
113        .await;
114}
115
116#[derive(Clone, Copy)]
117struct LocalExecutor;
118
119impl<F: Future<Output = ()> + 'static> Executor<F> for LocalExecutor {
120    fn execute(&self, future: F) {
121        tokio::task::spawn_local(future);
122    }
123}
124
125async fn handle_request(
126    router: Rc<Router>,
127    request: HyperRequest<Incoming>,
128    address: std::net::SocketAddr,
129) -> HyperResponse<HyperResponseBody> {
130    #[cfg(feature = "websocket")]
131    let mut request = request;
132    #[cfg(feature = "websocket")]
133    let on_upgrade = hyper::upgrade::on(&mut request);
134    let (parts, body) = request.into_parts();
135    let mut headers = Headers::new();
136
137    for (name, value) in &parts.headers {
138        headers
139            .append(name.as_str(), value.as_bytes())
140            .expect("Hyper supplied an invalid request header");
141    }
142
143    let mut request = Request::from_parts(
144        Method::try_from(parts.method.as_str()).expect("Hyper supplied an invalid request method"),
145        parts.uri.path(),
146        parts.uri.query().map(str::to_owned),
147        headers,
148        Box::new(HyperRequestStream::new(body)),
149    );
150    request.insert_extension(address);
151    let response = router.handle(request).await;
152
153    into_hyper_response(
154        response,
155        #[cfg(feature = "websocket")]
156        on_upgrade,
157    )
158}
159
160struct HyperRequestStream {
161    body: Pin<Box<Incoming>>,
162    current: Option<Bytes>,
163}
164
165impl HyperRequestStream {
166    fn new(body: Incoming) -> Self {
167        Self {
168            body: Box::pin(body),
169            current: None,
170        }
171    }
172}
173
174impl RequestStream for HyperRequestStream {
175    fn poll_next(&mut self, context: &mut Context<'_>) -> Poll<Option<Result<(), StreamError>>> {
176        loop {
177            match self.body.as_mut().poll_frame(context) {
178                Poll::Ready(Some(Ok(frame))) => match frame.into_data() {
179                    Ok(data) => {
180                        self.current = Some(data);
181                        return Poll::Ready(Some(Ok(())));
182                    }
183                    Err(_) => continue,
184                },
185                Poll::Ready(Some(Err(error))) => {
186                    self.current = None;
187                    return Poll::Ready(Some(Err(StreamError::new(error.to_string()))));
188                }
189                Poll::Ready(None) => {
190                    self.current = None;
191                    return Poll::Ready(None);
192                }
193                Poll::Pending => return Poll::Pending,
194            }
195        }
196    }
197
198    fn chunk(&self) -> &[u8] {
199        self.current.as_deref().unwrap_or_default()
200    }
201}
202
203struct HyperResponseBody {
204    body: ResponseBody,
205}
206
207struct HyperChunk(Chunk);
208
209impl hyper::body::Buf for HyperChunk {
210    fn remaining(&self) -> usize {
211        self.0.remaining()
212    }
213
214    fn chunk(&self) -> &[u8] {
215        self.0.bytes()
216    }
217
218    fn advance(&mut self, count: usize) {
219        self.0.advance(count);
220    }
221}
222
223impl HyperResponseBody {
224    fn new(body: ResponseBody) -> Self {
225        Self { body }
226    }
227}
228
229impl HyperBody for HyperResponseBody {
230    type Data = HyperChunk;
231    type Error = StreamError;
232
233    fn poll_frame(
234        self: Pin<&mut Self>,
235        context: &mut Context<'_>,
236    ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
237        match &mut self.get_mut().body {
238            ResponseBody::Buffered(bytes) if bytes.is_empty() => Poll::Ready(None),
239            ResponseBody::Buffered(bytes) => {
240                let chunk = HyperChunk(Chunk::from(std::mem::take(bytes)));
241                Poll::Ready(Some(Ok(Frame::data(chunk))))
242            }
243            ResponseBody::Streaming(stream) => match stream.poll_next(context) {
244                Poll::Ready(Some(Ok(chunk))) => {
245                    Poll::Ready(Some(Ok(Frame::data(HyperChunk(chunk)))))
246                }
247                Poll::Ready(Some(Err(error))) => Poll::Ready(Some(Err(error))),
248                Poll::Ready(None) => Poll::Ready(None),
249                Poll::Pending => Poll::Pending,
250            },
251            #[cfg(feature = "websocket")]
252            ResponseBody::WebSocket(_) => Poll::Ready(None),
253        }
254    }
255
256    fn is_end_stream(&self) -> bool {
257        if matches!(&self.body, ResponseBody::Buffered(bytes) if bytes.is_empty()) {
258            return true;
259        }
260
261        #[cfg(feature = "websocket")]
262        if matches!(&self.body, ResponseBody::WebSocket(_)) {
263            return true;
264        }
265
266        false
267    }
268
269    fn size_hint(&self) -> SizeHint {
270        let mut hint = SizeHint::new();
271
272        if let ResponseBody::Buffered(bytes) = &self.body {
273            hint.set_exact(bytes.len() as u64);
274        }
275
276        hint
277    }
278}
279
280fn into_hyper_response(
281    response: Response,
282    #[cfg(feature = "websocket")] on_upgrade: hyper::upgrade::OnUpgrade,
283) -> HyperResponse<HyperResponseBody> {
284    let (status, headers, body) = response.into_parts();
285
286    #[cfg(feature = "websocket")]
287    let body = match body {
288        ResponseBody::WebSocket(plan) => {
289            return into_hyper_websocket_response(headers, plan, on_upgrade);
290        }
291        body => body,
292    };
293
294    let length = body.buffered().map(<[u8]>::len);
295    let has_content_length = headers.contains("content-length");
296    let mut response = HyperResponse::new(HyperResponseBody::new(body));
297
298    match hyper::StatusCode::from_u16(status) {
299        Ok(status) => *response.status_mut() = status,
300        Err(_) => *response.status_mut() = hyper::StatusCode::INTERNAL_SERVER_ERROR,
301    }
302
303    append_headers(&mut response, headers);
304
305    if !has_content_length && let Some(length) = length {
306        response.headers_mut().insert(CONTENT_LENGTH, length.into());
307    }
308
309    response
310}
311
312fn append_headers(response: &mut HyperResponse<HyperResponseBody>, headers: Headers) {
313    for (name, value) in headers.iter() {
314        let name = HeaderName::from_bytes(name.as_bytes())
315            .expect("ServerKit generated an invalid response header name");
316        let value = HeaderValue::from_bytes(value)
317            .expect("ServerKit generated an invalid response header value");
318
319        response.headers_mut().append(name, value);
320    }
321}
322
323#[cfg(feature = "websocket")]
324fn into_hyper_websocket_response(
325    headers: Headers,
326    plan: WebSocketPlan,
327    on_upgrade: hyper::upgrade::OnUpgrade,
328) -> HyperResponse<HyperResponseBody> {
329    let accept_key = derive_accept_key(plan.key().as_bytes());
330
331    tokio::task::spawn_local(async move {
332        let Ok(upgraded) = on_upgrade.await else {
333            return;
334        };
335        let stream =
336            WebSocketStream::from_raw_socket(TokioIo::new(upgraded), Role::Server, None).await;
337        plan.run(WebSocket::from_io(HyperWebSocket { stream }))
338            .await;
339    });
340
341    let mut response =
342        HyperResponse::new(HyperResponseBody::new(ResponseBody::Buffered(Vec::new())));
343    *response.status_mut() = hyper::StatusCode::SWITCHING_PROTOCOLS;
344    response
345        .headers_mut()
346        .insert("connection", HeaderValue::from_static("Upgrade"));
347    response
348        .headers_mut()
349        .insert("upgrade", HeaderValue::from_static("websocket"));
350    response.headers_mut().insert(
351        "sec-websocket-accept",
352        HeaderValue::from_str(&accept_key)
353            .expect("a derived WebSocket accept key is a valid header value"),
354    );
355    append_headers(&mut response, headers);
356
357    response
358}
359
360#[cfg(feature = "websocket")]
361struct HyperWebSocket {
362    stream: WebSocketStream<TokioIo<hyper::upgrade::Upgraded>>,
363}
364
365#[cfg(feature = "websocket")]
366impl WebSocketIo for HyperWebSocket {
367    fn poll_next(
368        &mut self,
369        context: &mut Context<'_>,
370    ) -> Poll<Option<Result<WebSocketMessage, WebSocketError>>> {
371        loop {
372            let next = match Stream::poll_next(Pin::new(&mut self.stream), context) {
373                Poll::Ready(next) => next,
374                Poll::Pending => return Poll::Pending,
375            };
376
377            return Poll::Ready(match next {
378                Some(Ok(Message::Text(text))) => Some(Ok(WebSocketMessage::Text(text.to_string()))),
379                Some(Ok(Message::Binary(bytes))) => {
380                    Some(Ok(WebSocketMessage::Binary(bytes.to_vec())))
381                }
382                Some(Ok(Message::Ping(bytes))) => Some(Ok(WebSocketMessage::Ping(bytes.to_vec()))),
383                Some(Ok(Message::Pong(bytes))) => Some(Ok(WebSocketMessage::Pong(bytes.to_vec()))),
384                Some(Ok(Message::Close(frame))) => Some(Ok(WebSocketMessage::Close {
385                    code: frame.as_ref().map(|frame| u16::from(frame.code)),
386                    reason: frame.map_or_else(String::new, |frame| frame.reason.to_string()),
387                })),
388                Some(Ok(Message::Frame(_))) => continue,
389                Some(Err(error)) => Some(Err(WebSocketError::new(error.to_string()))),
390                None => None,
391            });
392        }
393    }
394
395    fn poll_ready(&mut self, context: &mut Context<'_>) -> Poll<Result<(), WebSocketError>> {
396        Sink::poll_ready(Pin::new(&mut self.stream), context)
397            .map_err(|error| WebSocketError::new(error.to_string()))
398    }
399
400    fn start_send(&mut self, message: WebSocketMessage) -> Result<(), WebSocketError> {
401        Sink::start_send(Pin::new(&mut self.stream), into_hyper_message(message))
402            .map_err(|error| WebSocketError::new(error.to_string()))
403    }
404
405    fn poll_flush(&mut self, context: &mut Context<'_>) -> Poll<Result<(), WebSocketError>> {
406        Sink::poll_flush(Pin::new(&mut self.stream), context)
407            .map_err(|error| WebSocketError::new(error.to_string()))
408    }
409
410    fn poll_close(&mut self, context: &mut Context<'_>) -> Poll<Result<(), WebSocketError>> {
411        Sink::poll_close(Pin::new(&mut self.stream), context)
412            .map_err(|error| WebSocketError::new(error.to_string()))
413    }
414}
415
416#[cfg(feature = "websocket")]
417fn into_hyper_message(message: WebSocketMessage) -> Message {
418    match message {
419        WebSocketMessage::Text(text) => Message::Text(text.into()),
420        WebSocketMessage::Binary(bytes) => Message::Binary(bytes.into()),
421        WebSocketMessage::Ping(bytes) => Message::Ping(bytes.into()),
422        WebSocketMessage::Pong(bytes) => Message::Pong(bytes.into()),
423        WebSocketMessage::Close { code, reason } => Message::Close(code.map(|code| CloseFrame {
424            code: CloseCode::from(code),
425            reason: reason.into(),
426        })),
427    }
428}
429
430#[cfg(test)]
431mod tests {
432    use std::{
433        net::TcpListener,
434        rc::Rc,
435        task::{Context, Poll},
436    };
437
438    use http_body_util::{BodyExt, Empty};
439    use hyper::{Request, body::Bytes, client::conn::http2};
440    use hyper_util::rt::TokioIo;
441    use serverkit::{Chunk, Config, Response, ResponseStream, RouteMethods, Router, StreamError};
442    use tokio::io::{AsyncReadExt, AsyncWriteExt};
443
444    use super::{LocalExecutor, serve_connection};
445
446    struct LargeStream {
447        sent: bool,
448    }
449
450    impl ResponseStream for LargeStream {
451        fn poll_next(
452            &mut self,
453            _context: &mut Context<'_>,
454        ) -> Poll<Option<Result<Chunk, StreamError>>> {
455            if self.sent {
456                Poll::Ready(None)
457            } else {
458                self.sent = true;
459                Poll::Ready(Some(Ok(Chunk::from(vec![7; 1024 * 1024]))))
460            }
461        }
462    }
463
464    fn router() -> Router {
465        Router::new(
466            Config::new(),
467            (
468                "/health".GET(|| async { "ok" }),
469                "/stream".GET(|| async { Response::stream(200, LargeStream { sent: false }) }),
470            ),
471        )
472    }
473
474    fn listener() -> (TcpListener, std::net::SocketAddr) {
475        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
476        let address = listener.local_addr().unwrap();
477        listener.set_nonblocking(true).unwrap();
478        (listener, address)
479    }
480
481    fn runtime() -> tokio::runtime::Runtime {
482        tokio::runtime::Builder::new_current_thread()
483            .enable_io()
484            .build()
485            .unwrap()
486    }
487
488    fn request_http_1(version: &str) -> String {
489        let (listener, address) = listener();
490        let runtime = runtime();
491
492        tokio::task::LocalSet::new().block_on(&runtime, async move {
493            let listener = tokio::net::TcpListener::from_std(listener).unwrap();
494            let server = tokio::task::spawn_local(async move {
495                let (connection, peer) = listener.accept().await.unwrap();
496                serve_connection(Rc::new(router()), connection, peer).await;
497            });
498            let mut client = tokio::net::TcpStream::connect(address).await.unwrap();
499            let request =
500                format!("GET /health {version}\r\nHost: localhost\r\nConnection: close\r\n\r\n");
501            client.write_all(request.as_bytes()).await.unwrap();
502            let mut response = Vec::new();
503            client.read_to_end(&mut response).await.unwrap();
504            server.await.unwrap();
505
506            String::from_utf8(response).unwrap()
507        })
508    }
509
510    #[test]
511    fn serves_http_1_0_and_http_1_1() {
512        assert!(request_http_1("HTTP/1.0").starts_with("HTTP/1.0 200 OK"));
513        assert!(request_http_1("HTTP/1.1").starts_with("HTTP/1.1 200 OK"));
514    }
515
516    #[test]
517    fn serves_http_2() {
518        let (listener, address) = listener();
519        let runtime = runtime();
520
521        tokio::task::LocalSet::new().block_on(&runtime, async move {
522            let listener = tokio::net::TcpListener::from_std(listener).unwrap();
523            tokio::task::spawn_local(async move {
524                let (connection, peer) = listener.accept().await.unwrap();
525                serve_connection(Rc::new(router()), connection, peer).await;
526            });
527            let client = tokio::net::TcpStream::connect(address).await.unwrap();
528            let (mut sender, connection) = http2::Builder::new(LocalExecutor)
529                .handshake(TokioIo::new(client))
530                .await
531                .unwrap();
532            tokio::task::spawn_local(async move {
533                connection.await.unwrap();
534            });
535            let request = Request::builder()
536                .uri("http://localhost/health")
537                .body(Empty::<Bytes>::new())
538                .unwrap();
539            let response = sender.send_request(request).await.unwrap();
540
541            assert_eq!(response.version(), hyper::Version::HTTP_2);
542            assert_eq!(response.status(), hyper::StatusCode::OK);
543
544            let request = Request::builder()
545                .uri("http://localhost/stream")
546                .body(Empty::<Bytes>::new())
547                .unwrap();
548            let response = sender.send_request(request).await.unwrap();
549            let body = response.into_body().collect().await.unwrap().to_bytes();
550
551            assert_eq!(body.len(), 1024 * 1024);
552            assert!(body.iter().all(|byte| *byte == 7));
553        });
554    }
555}