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}