1#[cfg(feature = "http2")]
8mod http2;
9
10use std::cell::RefCell;
11use std::collections::VecDeque;
12use std::fmt;
13use std::future::Future;
14use std::io::{self, Write as _};
15use std::net::{SocketAddr, ToSocketAddrs};
16use std::num::NonZeroUsize;
17use std::pin::Pin;
18use std::rc::Rc;
19#[cfg(feature = "tls")]
20use std::sync::Arc as TlsArc;
21use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
22use std::sync::{Arc, Mutex, Once};
23use std::task::{Context, Poll, Waker};
24use std::time::{Duration, Instant};
25
26use blazingly_core::{BackgroundTask, BodyStream, BodyStreamError, HttpMethod, HttpUpgrade};
27use blazingly_core::{StreamingBody, UpgradeIoError, UpgradedIo};
28use blazingly_executor::{ExecutableApp, InvocationControl};
29pub use blazingly_http::HttpMiddleware;
30use blazingly_http::{HttpApp, HttpRequestView, Response};
31use blazingly_openapi::OpenApiConfig;
32use blazingly_wire::{BodyFraming, ChunkDecoder, HeaderPositions, reason_phrase};
33use blazingly_wire::{StreamingChunk, StreamingChunkDecoder};
34use compio::dispatcher::Dispatcher;
35#[cfg(any(feature = "http2", feature = "tls"))]
36use compio::io::compat::AsyncStream;
37use compio::io::{AsyncReadExt as CompioAsyncReadExt, AsyncWriteExt as CompioAsyncWriteExt};
38use compio::net::TcpListener;
39use compio::net::TcpStream;
40use compio::runtime::{Runtime, spawn};
41use futures_lite::future;
42use futures_lite::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
43
44pub const DEFAULT_MAX_BODY_BYTES: usize = blazingly_wire::DEFAULT_MAX_BODY_BYTES;
46
47const DEFAULT_MAX_HEADER_BYTES: usize = blazingly_wire::DEFAULT_MAX_HEADER_BYTES;
48const DEFAULT_MAX_HEADERS: usize = blazingly_wire::DEFAULT_MAX_HEADERS;
49const DEFAULT_MAX_CHUNKS: usize = blazingly_wire::DEFAULT_MAX_CHUNKS;
50const DEFAULT_MAX_PIPELINE_BATCH: usize = 16;
51const CONTINUE_RESPONSE: &[u8] = b"HTTP/1.1 100 Continue\r\n\r\n";
52const MAX_PIPELINE_WRITE_BYTES: usize = 64 * 1024;
53const MAX_HEADER_CAPACITY: usize = blazingly_wire::MAX_HEADER_CAPACITY;
54const READ_CHUNK_BYTES: usize = 8 * 1024;
55const DEFAULT_STREAM_READ_BYTES: usize = 64 * 1024;
56const MAX_SPARE_BUFFERS: usize = 4;
61const DEFAULT_DRAIN_TIMEOUT: Duration = Duration::from_secs(30);
62const DEFAULT_HEADER_READ_TIMEOUT: Duration = Duration::from_secs(10);
63const DEFAULT_BODY_READ_TIMEOUT: Duration = Duration::from_secs(30);
64const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(75);
65const DEFAULT_WRITE_TIMEOUT: Duration = Duration::from_secs(30);
66static NEXT_MULTICORE_SERVER_ID: AtomicU64 = AtomicU64::new(1);
67static DATE_UPDATER: Once = Once::new();
68static DATE_GENERATION: AtomicU64 = AtomicU64::new(0);
69static DATE_VALUE: Mutex<String> = Mutex::new(String::new());
70
71thread_local! {
72 static WORKER_APPS: RefCell<std::collections::HashMap<u64, Rc<HttpApp>>> =
73 RefCell::new(std::collections::HashMap::new());
74 static DRAIN_COUNTER: RefCell<Option<Arc<AtomicUsize>>> = const { RefCell::new(None) };
75 static DATE_CACHE: RefCell<CachedDate> = const {
76 RefCell::new(CachedDate {
77 generation: u64::MAX,
78 value: String::new(),
79 })
80 };
81}
82
83struct CachedDate {
84 generation: u64,
85 value: String,
86}
87
88#[derive(Clone, Copy, Debug, Eq, PartialEq)]
97#[allow(clippy::struct_field_names)]
98pub struct ServerLimits {
99 max_header_bytes: usize,
100 max_headers: usize,
101 max_body_bytes: usize,
102 max_chunks: usize,
103 max_pipeline_batch: usize,
104 max_requests_per_connection: Option<NonZeroUsize>,
105 stream_read_bytes: usize,
106 header_read_timeout: Duration,
107 body_read_timeout: Duration,
108 idle_timeout: Duration,
109 write_timeout: Duration,
110 #[cfg(feature = "http2")]
111 max_concurrent_streams: usize,
112}
113
114impl ServerLimits {
115 #[must_use]
116 pub const fn new() -> Self {
117 Self {
118 max_header_bytes: DEFAULT_MAX_HEADER_BYTES,
119 max_headers: DEFAULT_MAX_HEADERS,
120 max_body_bytes: DEFAULT_MAX_BODY_BYTES,
121 max_chunks: DEFAULT_MAX_CHUNKS,
122 max_pipeline_batch: DEFAULT_MAX_PIPELINE_BATCH,
123 max_requests_per_connection: None,
124 stream_read_bytes: DEFAULT_STREAM_READ_BYTES,
125 header_read_timeout: DEFAULT_HEADER_READ_TIMEOUT,
126 body_read_timeout: DEFAULT_BODY_READ_TIMEOUT,
127 idle_timeout: DEFAULT_IDLE_TIMEOUT,
128 write_timeout: DEFAULT_WRITE_TIMEOUT,
129 #[cfg(feature = "http2")]
130 max_concurrent_streams: 100,
131 }
132 }
133
134 #[must_use]
135 pub const fn with_max_header_bytes(mut self, bytes: usize) -> Self {
136 assert!(bytes > 0, "max_header_bytes must be greater than zero");
137 self.max_header_bytes = bytes;
138 self
139 }
140
141 #[must_use]
146 pub const fn with_max_headers(mut self, count: usize) -> Self {
147 assert!(count > 0, "max_headers must be greater than zero");
148 assert!(
149 count <= MAX_HEADER_CAPACITY,
150 "max_headers cannot exceed the native stack capacity"
151 );
152 self.max_headers = count;
153 self
154 }
155
156 #[must_use]
157 pub const fn with_max_body_bytes(mut self, bytes: usize) -> Self {
158 self.max_body_bytes = bytes;
159 self
160 }
161
162 #[must_use]
164 pub const fn with_max_chunks(mut self, count: usize) -> Self {
165 assert!(count > 0, "max_chunks must be greater than zero");
166 self.max_chunks = count;
167 self
168 }
169
170 #[must_use]
173 pub const fn with_max_pipeline_batch(mut self, count: NonZeroUsize) -> Self {
174 self.max_pipeline_batch = count.get();
175 self
176 }
177
178 #[must_use]
180 pub const fn with_max_requests_per_connection(mut self, count: Option<NonZeroUsize>) -> Self {
181 self.max_requests_per_connection = count;
182 self
183 }
184
185 #[must_use]
199 pub const fn with_stream_read_bytes(mut self, bytes: usize) -> Self {
200 assert!(bytes > 0, "stream_read_bytes must be greater than zero");
201 self.stream_read_bytes = bytes;
202 self
203 }
204
205 #[must_use]
208 pub const fn with_header_read_timeout(mut self, timeout: Duration) -> Self {
209 self.header_read_timeout = timeout;
210 self
211 }
212
213 #[must_use]
215 pub const fn with_body_read_timeout(mut self, timeout: Duration) -> Self {
216 self.body_read_timeout = timeout;
217 self
218 }
219
220 #[must_use]
223 pub const fn with_idle_timeout(mut self, timeout: Duration) -> Self {
224 self.idle_timeout = timeout;
225 self
226 }
227
228 #[must_use]
230 pub const fn with_write_timeout(mut self, timeout: Duration) -> Self {
231 self.write_timeout = timeout;
232 self
233 }
234
235 #[cfg(feature = "http2")]
236 #[must_use]
237 pub const fn with_max_concurrent_streams(mut self, count: NonZeroUsize) -> Self {
238 self.max_concurrent_streams = count.get();
239 self
240 }
241
242 #[must_use]
243 pub const fn max_header_bytes(self) -> usize {
244 self.max_header_bytes
245 }
246
247 #[must_use]
248 pub const fn max_headers(self) -> usize {
249 self.max_headers
250 }
251
252 #[must_use]
253 pub const fn max_body_bytes(self) -> usize {
254 self.max_body_bytes
255 }
256
257 #[must_use]
258 pub const fn max_chunks(self) -> usize {
259 self.max_chunks
260 }
261
262 #[must_use]
263 pub const fn max_pipeline_batch(self) -> usize {
264 self.max_pipeline_batch
265 }
266
267 #[must_use]
268 pub const fn max_requests_per_connection(self) -> Option<NonZeroUsize> {
269 self.max_requests_per_connection
270 }
271
272 #[must_use]
273 pub const fn stream_read_bytes(self) -> usize {
274 self.stream_read_bytes
275 }
276
277 #[must_use]
278 pub const fn header_read_timeout(self) -> Duration {
279 self.header_read_timeout
280 }
281
282 #[must_use]
283 pub const fn body_read_timeout(self) -> Duration {
284 self.body_read_timeout
285 }
286
287 #[must_use]
288 pub const fn idle_timeout(self) -> Duration {
289 self.idle_timeout
290 }
291
292 #[must_use]
293 pub const fn write_timeout(self) -> Duration {
294 self.write_timeout
295 }
296
297 #[cfg(feature = "http2")]
298 #[must_use]
299 pub const fn max_concurrent_streams(self) -> usize {
300 self.max_concurrent_streams
301 }
302}
303
304impl Default for ServerLimits {
305 fn default() -> Self {
306 Self::new()
307 }
308}
309
310#[derive(Clone)]
312pub struct ShutdownHandle {
313 state: Arc<ShutdownState>,
314}
315
316pub struct ShutdownSignal {
318 state: Arc<ShutdownState>,
319}
320
321struct ShutdownState {
322 requested: AtomicBool,
323 waker: Mutex<Option<Waker>>,
324}
325
326#[must_use]
328pub fn shutdown_channel() -> (ShutdownHandle, ShutdownSignal) {
329 let state = Arc::new(ShutdownState {
330 requested: AtomicBool::new(false),
331 waker: Mutex::new(None),
332 });
333 (
334 ShutdownHandle {
335 state: Arc::clone(&state),
336 },
337 ShutdownSignal { state },
338 )
339}
340
341pub fn termination_channel() -> io::Result<(ShutdownHandle, ShutdownSignal)> {
352 let (handle, signal) = shutdown_channel();
353 let termination = handle.clone();
354 ctrlc::set_handler(move || termination.shutdown())
355 .map_err(|error| io::Error::other(error.to_string()))?;
356 Ok((handle, signal))
357}
358
359impl ShutdownHandle {
360 pub fn shutdown(&self) {
361 self.state.requested.store(true, Ordering::Release);
362 if let Some(waker) = self
363 .state
364 .waker
365 .lock()
366 .unwrap_or_else(std::sync::PoisonError::into_inner)
367 .take()
368 {
369 waker.wake();
370 }
371 }
372
373 #[must_use]
374 pub fn is_shutdown(&self) -> bool {
375 self.state.requested.load(Ordering::Acquire)
376 }
377}
378
379impl Future for ShutdownSignal {
380 type Output = ();
381
382 fn poll(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
383 if self.state.requested.load(Ordering::Acquire) {
384 return Poll::Ready(());
385 }
386 *self
387 .state
388 .waker
389 .lock()
390 .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(context.waker().clone());
391 if self.state.requested.load(Ordering::Acquire) {
392 Poll::Ready(())
393 } else {
394 Poll::Pending
395 }
396 }
397}
398
399pub struct Server {
404 app: Rc<HttpApp>,
405 limits: ServerLimits,
406 request_timeout: Option<Duration>,
407 max_connections: Option<NonZeroUsize>,
408 tcp_nodelay: bool,
409 #[cfg(feature = "tls")]
410 tls_acceptor: Option<compio::tls::TlsAcceptor>,
411}
412
413impl Server {
414 #[must_use]
415 pub fn new(app: ExecutableApp) -> Self {
416 Self {
417 app: Rc::new(HttpApp::new(app)),
418 limits: ServerLimits::new(),
419 request_timeout: None,
420 max_connections: None,
421 tcp_nodelay: true,
422 #[cfg(feature = "tls")]
423 tls_acceptor: None,
424 }
425 }
426
427 #[must_use]
428 pub fn with_max_body_bytes(mut self, max_body_bytes: usize) -> Self {
429 self.app = Rc::new(
430 Rc::try_unwrap(self.app)
431 .unwrap_or_else(|_| unreachable!("server app is not shared before serving"))
432 .with_max_body_bytes(max_body_bytes),
433 );
434 self.limits = self.limits.with_max_body_bytes(max_body_bytes);
435 self
436 }
437
438 #[must_use]
444 pub fn with_middleware(mut self, middleware: impl HttpMiddleware + 'static) -> Self {
445 self.app = Rc::new(
446 Rc::try_unwrap(self.app)
447 .unwrap_or_else(|_| unreachable!("server app is not shared before serving"))
448 .with_middleware(middleware),
449 );
450 self
451 }
452
453 #[must_use]
455 pub fn with_shared_middleware(mut self, middleware: Rc<dyn HttpMiddleware>) -> Self {
456 self.app = Rc::new(
457 Rc::try_unwrap(self.app)
458 .unwrap_or_else(|_| unreachable!("server app is not shared before serving"))
459 .with_shared_middleware(middleware),
460 );
461 self
462 }
463
464 #[must_use]
469 pub const fn with_max_connections(mut self, connections: Option<NonZeroUsize>) -> Self {
470 self.max_connections = connections;
471 self
472 }
473
474 #[must_use]
477 pub const fn with_tcp_nodelay(mut self, nodelay: bool) -> Self {
478 self.tcp_nodelay = nodelay;
479 self
480 }
481
482 #[must_use]
483 pub fn with_limits(mut self, limits: ServerLimits) -> Self {
484 self.app = Rc::new(
485 Rc::try_unwrap(self.app)
486 .unwrap_or_else(|_| unreachable!("server app is not shared before serving"))
487 .with_max_body_bytes(limits.max_body_bytes),
488 );
489 self.limits = limits;
490 self
491 }
492
493 #[must_use]
495 pub fn with_openapi(mut self, config: OpenApiConfig) -> Self {
496 self.app = Rc::new(
497 Rc::try_unwrap(self.app)
498 .unwrap_or_else(|_| unreachable!("server app is not shared before serving"))
499 .with_openapi(config),
500 );
501 self
502 }
503
504 #[must_use]
509 pub const fn with_request_timeout(mut self, timeout: Duration) -> Self {
510 self.request_timeout = Some(timeout);
511 self
512 }
513
514 #[cfg(feature = "tls")]
519 #[must_use]
520 pub fn with_tls_config(mut self, config: TlsArc<compio::tls::rustls::ServerConfig>) -> Self {
521 self.tls_acceptor = Some(config.into());
522 self
523 }
524
525 pub fn serve(self, address: impl ToSocketAddrs) -> io::Result<()> {
532 let address = resolve_address(address)?;
533 let runtime = Runtime::new()?;
534 runtime.block_on(async {
535 let listener = TcpListener::bind(address).await?;
536 self.serve_listener(listener, None, DEFAULT_DRAIN_TIMEOUT)
537 .await
538 })
539 }
540
541 pub fn serve_gracefully(
548 self,
549 address: impl ToSocketAddrs,
550 shutdown: ShutdownSignal,
551 drain_timeout: Duration,
552 ) -> io::Result<()> {
553 let address = resolve_address(address)?;
554 let runtime = Runtime::new()?;
555 runtime.block_on(async {
556 let listener = TcpListener::bind(address).await?;
557 self.serve_listener(listener, Some(shutdown), drain_timeout)
558 .await
559 })
560 }
561
562 pub async fn serve_io<IO>(&self, io: &mut IO) -> io::Result<()>
572 where
573 IO: AsyncRead + AsyncWrite + Unpin,
574 {
575 serve_connection(
576 self.app.as_ref(),
577 self.limits,
578 io,
579 None,
580 "http",
581 None,
582 self.request_timeout,
583 TransportOwnership::Borrowed,
584 )
585 .await
586 .map(|_| ())
587 }
588
589 pub async fn serve_owned_io<IO>(&self, io: IO) -> io::Result<()>
600 where
601 IO: AsyncRead + AsyncWrite + Unpin + 'static,
602 {
603 serve_compat_connection(
604 self.app.as_ref(),
605 self.limits,
606 io,
607 None,
608 "http",
609 None,
610 self.request_timeout,
611 )
612 .await
613 }
614
615 #[cfg(feature = "http2")]
626 pub async fn serve_http2_io<IO>(&self, io: &mut IO) -> io::Result<()>
627 where
628 IO: AsyncRead + AsyncWrite + Unpin,
629 {
630 http2::serve_connection(
631 self.app.as_ref(),
632 self.limits,
633 io,
634 None,
635 "http",
636 None,
637 Vec::new(),
638 self.request_timeout,
639 )
640 .await
641 }
642
643 fn connection_setup(&self) -> ConnectionSetup {
644 ConnectionSetup {
645 limits: self.limits,
646 request_timeout: self.request_timeout,
647 #[cfg(feature = "tls")]
648 tls_acceptor: self.tls_acceptor.clone(),
649 }
650 }
651
652 async fn serve_listener(
653 &self,
654 listener: TcpListener,
655 shutdown: Option<ShutdownSignal>,
656 drain_timeout: Duration,
657 ) -> io::Result<()> {
658 self.app
659 .startup()
660 .await
661 .map_err(|error| io::Error::other(error.to_string()))?;
662 let active = Arc::new(AtomicUsize::new(0));
663 let connections = Arc::new(AtomicUsize::new(0));
664 let _drain_scope = DrainScope::new(&active);
665 let shutdown_state = shutdown
666 .as_ref()
667 .map(|shutdown| Arc::clone(&shutdown.state));
668 let mut shutdown = shutdown.map(Box::pin);
669 loop {
670 let accepted = if let Some(shutdown) = shutdown.as_mut() {
671 future::race(
672 async { AcceptEvent::Connection(listener.accept().await) },
673 async {
674 shutdown.as_mut().await;
675 AcceptEvent::Shutdown
676 },
677 )
678 .await
679 } else {
680 AcceptEvent::Connection(listener.accept().await)
681 };
682 let (stream, peer_addr) = match accepted {
683 AcceptEvent::Connection(result) => result?,
684 AcceptEvent::Shutdown => break,
685 };
686 if self
687 .max_connections
688 .is_some_and(|limit| connections.load(Ordering::Acquire) >= limit.get())
689 {
690 drop(stream);
691 continue;
692 }
693 let _ = stream.set_nodelay(self.tcp_nodelay);
695 let app = Rc::clone(&self.app);
696 let setup = self.connection_setup();
697 active.fetch_add(1, Ordering::AcqRel);
698 connections.fetch_add(1, Ordering::AcqRel);
699 let active_for_task = Arc::clone(&active);
700 let connections_for_task = Arc::clone(&connections);
701 let connection_shutdown = shutdown_state.clone();
702 spawn(async move {
703 let _work = ActiveWork::adopt(active_for_task);
704 let _slot = ActiveWork::adopt(connections_for_task);
705 serve_accepted(
706 app.as_ref(),
707 stream,
708 peer_addr,
709 connection_shutdown.as_deref(),
710 setup,
711 )
712 .await;
713 })
714 .detach();
715 }
716 let drain = async {
717 while active.load(Ordering::Acquire) != 0 {
718 compio::time::sleep(Duration::from_millis(5)).await;
719 }
720 };
721 let _ = compio::time::timeout(drain_timeout, drain).await;
722 self.app
723 .shutdown()
724 .await
725 .map_err(|error| io::Error::other(error.to_string()))
726 }
727}
728
729pub struct MulticoreServer<Factory> {
734 factory: Factory,
735 workers: NonZeroUsize,
736 limits: ServerLimits,
737 request_timeout: Option<Duration>,
738 openapi: Option<OpenApiConfig>,
739 middleware: Option<MiddlewareFactory>,
740 max_connections: Option<NonZeroUsize>,
741 tcp_nodelay: bool,
742 #[cfg(feature = "tls")]
743 tls_acceptor: Option<compio::tls::TlsAcceptor>,
744}
745
746impl<Factory> MulticoreServer<Factory>
747where
748 Factory: Fn() -> ExecutableApp + Send + Sync + 'static,
749{
750 #[must_use]
751 pub fn new(workers: NonZeroUsize, factory: Factory) -> Self {
752 Self {
753 factory,
754 workers,
755 limits: ServerLimits::new(),
756 request_timeout: None,
757 openapi: None,
758 middleware: None,
759 max_connections: None,
760 tcp_nodelay: true,
761 #[cfg(feature = "tls")]
762 tls_acceptor: None,
763 }
764 }
765
766 #[must_use]
767 pub const fn with_max_body_bytes(mut self, max_body_bytes: usize) -> Self {
768 self.limits = self.limits.with_max_body_bytes(max_body_bytes);
769 self
770 }
771
772 #[must_use]
773 pub const fn with_limits(mut self, limits: ServerLimits) -> Self {
774 self.limits = limits;
775 self
776 }
777
778 #[must_use]
780 pub const fn with_request_timeout(mut self, timeout: Duration) -> Self {
781 self.request_timeout = Some(timeout);
782 self
783 }
784
785 #[must_use]
787 pub fn with_openapi(mut self, config: OpenApiConfig) -> Self {
788 self.openapi = Some(config);
789 self
790 }
791
792 #[must_use]
798 pub fn with_middleware_factory<Middleware>(mut self, middleware: Middleware) -> Self
799 where
800 Middleware: Fn() -> Vec<Rc<dyn HttpMiddleware>> + Send + Sync + 'static,
801 {
802 self.middleware = Some(Arc::new(middleware));
803 self
804 }
805
806 #[must_use]
811 pub const fn with_max_connections(mut self, connections: Option<NonZeroUsize>) -> Self {
812 self.max_connections = connections;
813 self
814 }
815
816 #[must_use]
819 pub const fn with_tcp_nodelay(mut self, nodelay: bool) -> Self {
820 self.tcp_nodelay = nodelay;
821 self
822 }
823
824 #[cfg(feature = "tls")]
825 #[must_use]
826 pub fn with_tls_config(mut self, config: TlsArc<compio::tls::rustls::ServerConfig>) -> Self {
827 self.tls_acceptor = Some(config.into());
828 self
829 }
830
831 pub fn serve(self, address: impl ToSocketAddrs) -> io::Result<()> {
842 self.serve_inner(address, None, DEFAULT_DRAIN_TIMEOUT)
843 }
844
845 #[allow(clippy::needless_pass_by_value)]
853 pub fn serve_gracefully(
854 self,
855 address: impl ToSocketAddrs,
856 shutdown: ShutdownSignal,
857 drain_timeout: Duration,
858 ) -> io::Result<()> {
859 self.serve_inner(address, Some(&shutdown.state), drain_timeout)
860 }
861
862 #[allow(clippy::too_many_lines)]
863 fn serve_inner(
864 self,
865 address: impl ToSocketAddrs,
866 shutdown: Option<&Arc<ShutdownState>>,
867 drain_timeout: Duration,
868 ) -> io::Result<()> {
869 let address = resolve_address(address)?;
870 let listener = std::net::TcpListener::bind(address)?;
871 listener.set_nonblocking(shutdown.is_some())?;
872 let dispatchers = (0..self.workers.get())
876 .map(|index| {
877 Dispatcher::builder()
878 .worker_threads(NonZeroUsize::new(1).expect("one worker is always non-zero"))
879 .thread_names(move |_| format!("blazingly-worker-{index}"))
880 .build()
881 })
882 .collect::<io::Result<Vec<_>>>()?;
883 let limits = self.limits;
884 let request_timeout = self.request_timeout;
885 let max_connections = self.max_connections;
886 let tcp_nodelay = self.tcp_nodelay;
887 #[cfg(feature = "tls")]
888 let tls_acceptor = self.tls_acceptor;
889 let config = Arc::new(WorkerConfig {
890 factory: self.factory,
891 max_body_bytes: limits.max_body_bytes,
892 openapi: self.openapi,
893 middleware: self.middleware,
894 });
895 let active = Arc::new(AtomicUsize::new(0));
896 let connections = Arc::new(AtomicUsize::new(0));
897 let server_id = NEXT_MULTICORE_SERVER_ID.fetch_add(1, Ordering::Relaxed);
898 let mut next_worker = 0_usize;
899
900 if let Err(error) = start_workers(&dispatchers, server_id, &config, &active) {
901 for dispatcher in dispatchers {
902 let _ = future::block_on(dispatcher.join());
903 }
904 return Err(error);
905 }
906
907 loop {
908 if shutdown
909 .as_ref()
910 .is_some_and(|shutdown| shutdown.requested.load(Ordering::Acquire))
911 {
912 break;
913 }
914 let (stream, peer_addr) = match listener.accept() {
915 Ok(connection) => connection,
916 Err(error) if error.kind() == io::ErrorKind::Interrupted => continue,
917 Err(error) if shutdown.is_some() && error.kind() == io::ErrorKind::WouldBlock => {
918 std::thread::sleep(Duration::from_millis(1));
919 continue;
920 }
921 Err(error) => return Err(error),
922 };
923
924 if max_connections
925 .is_some_and(|limit| connections.load(Ordering::Acquire) >= limit.get())
926 {
927 drop(stream);
928 continue;
929 }
930 let _ = stream.set_nodelay(tcp_nodelay);
932 let config = Arc::clone(&config);
933 let active_for_task = Arc::clone(&active);
934 let connections_for_task = Arc::clone(&connections);
935 let connection_shutdown = shutdown.cloned();
936 let setup = ConnectionSetup {
937 limits,
938 request_timeout,
939 #[cfg(feature = "tls")]
940 tls_acceptor: tls_acceptor.clone(),
941 };
942 active.fetch_add(1, Ordering::AcqRel);
943 connections.fetch_add(1, Ordering::AcqRel);
944 let dispatcher = &dispatchers[next_worker];
945 next_worker = (next_worker + 1) % dispatchers.len();
946 if dispatcher
947 .dispatch(move || async move {
948 let _work = ActiveWork::adopt(active_for_task);
949 let _slot = ActiveWork::adopt(connections_for_task);
950 let app = worker_app(server_id, config.as_ref()).0;
951 let Ok(stream) = compio::net::TcpStream::from_std(stream) else {
952 return;
953 };
954 serve_accepted(
955 app.as_ref(),
956 stream,
957 peer_addr,
958 connection_shutdown.as_deref(),
959 setup,
960 )
961 .await;
962 })
963 .is_err()
964 {
965 active.fetch_sub(1, Ordering::AcqRel);
966 connections.fetch_sub(1, Ordering::AcqRel);
967 return Err(io::Error::new(
968 io::ErrorKind::BrokenPipe,
969 "all Compio workers stopped",
970 ));
971 }
972 }
973
974 let deadline = Instant::now() + drain_timeout;
975 while active.load(Ordering::Acquire) != 0 && Instant::now() < deadline {
976 std::thread::sleep(Duration::from_millis(2));
977 }
978 let mut shutdown_tasks = Vec::with_capacity(dispatchers.len());
979 for dispatcher in &dispatchers {
980 let completion = dispatcher
981 .dispatch(move || async move {
982 clear_drain_counter();
983 if let Some(app) = take_worker_app(server_id) {
984 let _ = app.shutdown().await;
985 }
986 })
987 .map_err(|_| {
988 io::Error::new(
989 io::ErrorKind::BrokenPipe,
990 "Compio worker stopped before application shutdown",
991 )
992 })?;
993 shutdown_tasks.push(completion);
994 }
995 for completion in shutdown_tasks {
996 future::block_on(completion).map_err(|_| {
997 io::Error::new(
998 io::ErrorKind::BrokenPipe,
999 "Compio worker did not complete application shutdown",
1000 )
1001 })?;
1002 }
1003 for dispatcher in dispatchers {
1004 future::block_on(dispatcher.join())?;
1005 }
1006 Ok(())
1007 }
1008}
1009
1010type MiddlewareFactory = Arc<dyn Fn() -> Vec<Rc<dyn HttpMiddleware>> + Send + Sync>;
1012
1013struct WorkerConfig<Factory> {
1015 factory: Factory,
1016 max_body_bytes: usize,
1017 openapi: Option<OpenApiConfig>,
1018 middleware: Option<MiddlewareFactory>,
1019}
1020
1021fn start_workers<Factory>(
1023 dispatchers: &[Dispatcher],
1024 server_id: u64,
1025 config: &Arc<WorkerConfig<Factory>>,
1026 active: &Arc<AtomicUsize>,
1027) -> io::Result<()>
1028where
1029 Factory: Fn() -> ExecutableApp + Send + Sync + 'static,
1030{
1031 let mut started = Vec::with_capacity(dispatchers.len());
1032 for dispatcher in dispatchers {
1033 let config = Arc::clone(config);
1034 let active = Arc::clone(active);
1035 let completion = dispatcher
1036 .dispatch(move || async move {
1037 install_drain_counter(&active);
1038 let (app, created) = worker_app(server_id, config.as_ref());
1039 if created {
1040 return app.startup().await.map_err(|error| error.to_string());
1041 }
1042 Ok(())
1043 })
1044 .map_err(|_| {
1045 io::Error::new(
1046 io::ErrorKind::BrokenPipe,
1047 "Compio worker stopped before application startup",
1048 )
1049 })?;
1050 started.push(completion);
1051 }
1052 for completion in started {
1053 future::block_on(completion)
1054 .map_err(|_| {
1055 io::Error::new(
1056 io::ErrorKind::BrokenPipe,
1057 "Compio worker did not complete application startup",
1058 )
1059 })?
1060 .map_err(io::Error::other)?;
1061 }
1062 Ok(())
1063}
1064
1065fn worker_app<Factory>(server_id: u64, config: &WorkerConfig<Factory>) -> (Rc<HttpApp>, bool)
1066where
1067 Factory: Fn() -> ExecutableApp,
1068{
1069 WORKER_APPS.with(|apps| {
1070 let mut apps = apps.borrow_mut();
1071 if let Some(app) = apps.get(&server_id) {
1072 return (Rc::clone(app), false);
1073 }
1074 let app = HttpApp::new((config.factory)()).with_max_body_bytes(config.max_body_bytes);
1075 let app = match &config.openapi {
1076 Some(openapi) => app.with_openapi(openapi.clone()),
1077 None => app,
1078 };
1079 let app = match &config.middleware {
1080 Some(middleware) => middleware()
1081 .into_iter()
1082 .fold(app, HttpApp::with_shared_middleware),
1083 None => app,
1084 };
1085 let app = Rc::new(app);
1086 apps.insert(server_id, Rc::clone(&app));
1087 (app, true)
1088 })
1089}
1090
1091fn take_worker_app(server_id: u64) -> Option<Rc<HttpApp>> {
1092 WORKER_APPS.with(|apps| apps.borrow_mut().remove(&server_id))
1093}
1094
1095fn install_drain_counter(active: &Arc<AtomicUsize>) {
1096 DRAIN_COUNTER.with(|counter| *counter.borrow_mut() = Some(Arc::clone(active)));
1097}
1098
1099fn clear_drain_counter() {
1100 DRAIN_COUNTER.with(|counter| counter.borrow_mut().take());
1101}
1102
1103fn drain_counter() -> Option<Arc<AtomicUsize>> {
1104 DRAIN_COUNTER.with(|counter| counter.borrow().clone())
1105}
1106
1107struct DrainScope;
1110
1111impl DrainScope {
1112 fn new(active: &Arc<AtomicUsize>) -> Self {
1113 install_drain_counter(active);
1114 Self
1115 }
1116}
1117
1118impl Drop for DrainScope {
1119 fn drop(&mut self) {
1120 clear_drain_counter();
1121 }
1122}
1123
1124struct ActiveWork {
1126 active: Arc<AtomicUsize>,
1127}
1128
1129impl ActiveWork {
1130 fn acquire(active: Arc<AtomicUsize>) -> Self {
1132 active.fetch_add(1, Ordering::AcqRel);
1133 Self { active }
1134 }
1135
1136 const fn adopt(active: Arc<AtomicUsize>) -> Self {
1138 Self { active }
1139 }
1140}
1141
1142impl Drop for ActiveWork {
1143 fn drop(&mut self) {
1144 self.active.fetch_sub(1, Ordering::AcqRel);
1145 }
1146}
1147
1148enum AcceptEvent {
1149 Connection(io::Result<(compio::net::TcpStream, SocketAddr)>),
1150 Shutdown,
1151}
1152
1153struct ConnectionSetup {
1155 limits: ServerLimits,
1156 request_timeout: Option<Duration>,
1157 #[cfg(feature = "tls")]
1158 tls_acceptor: Option<compio::tls::TlsAcceptor>,
1159}
1160
1161async fn serve_accepted(
1167 app: &HttpApp,
1168 stream: compio::net::TcpStream,
1169 peer_addr: SocketAddr,
1170 shutdown: Option<&ShutdownState>,
1171 setup: ConnectionSetup,
1172) {
1173 let limits = setup.limits;
1174 let request_timeout = setup.request_timeout;
1175 #[cfg(feature = "tls")]
1176 if let Some(tls_acceptor) = setup.tls_acceptor {
1177 match within(limits.header_read_timeout, tls_acceptor.accept(stream)).await {
1178 Ok(Ok(stream)) => {
1179 let _ = serve_compat_connection(
1180 app,
1181 limits,
1182 Box::pin(AsyncStream::new(stream)),
1183 Some(peer_addr),
1184 "https",
1185 shutdown,
1186 request_timeout,
1187 )
1188 .await;
1189 }
1190 Ok(Err(error)) => report_failure("TLS handshake failed", &error),
1191 Err(Expired) => report_failure("TLS handshake", &"deadline expired"),
1192 }
1193 return;
1194 }
1195 let _ = serve_native_connection(
1196 app,
1197 limits,
1198 stream,
1199 Some(peer_addr),
1200 "http",
1201 shutdown,
1202 request_timeout,
1203 )
1204 .await;
1205}
1206
1207#[derive(Clone, Copy, Eq, PartialEq)]
1209enum TransportOwnership {
1210 Owned,
1212 Borrowed,
1214}
1215
1216enum ConnectionOutcome {
1218 Completed,
1220 Upgraded {
1225 upgrade: HttpUpgrade,
1226 buffered: Vec<u8>,
1227 },
1228}
1229
1230async fn serve_compat_connection<IO>(
1236 app: &HttpApp,
1237 limits: ServerLimits,
1238 mut io: IO,
1239 peer_addr: Option<SocketAddr>,
1240 scheme: &'static str,
1241 shutdown: Option<&ShutdownState>,
1242 request_timeout: Option<Duration>,
1243) -> io::Result<()>
1244where
1245 IO: AsyncRead + AsyncWrite + Unpin + 'static,
1246{
1247 let outcome = serve_connection(
1248 app,
1249 limits,
1250 &mut io,
1251 peer_addr,
1252 scheme,
1253 shutdown,
1254 request_timeout,
1255 TransportOwnership::Owned,
1256 )
1257 .await?;
1258 match outcome {
1259 ConnectionOutcome::Completed => Ok(()),
1260 ConnectionOutcome::Upgraded { upgrade, buffered } => upgrade
1261 .run(Box::new(CompatUpgradedIo { io, buffered }))
1262 .await
1263 .map_err(|error| io::Error::other(error.to_string())),
1264 }
1265}
1266
1267struct Expired;
1269
1270async fn within<Operation>(
1288 deadline: Duration,
1289 operation: Operation,
1290) -> Result<Operation::Output, Expired>
1291where
1292 Operation: Future,
1293{
1294 if compio::runtime::Runtime::try_with_current(|_| ()).is_err() {
1295 return Ok(operation.await);
1296 }
1297 let mut operation = std::pin::pin!(operation);
1298 let mut expiry = std::pin::pin!(compio::time::sleep(deadline));
1299 std::future::poll_fn(move |context| {
1300 if let Poll::Ready(output) = operation.as_mut().poll(context) {
1301 return Poll::Ready(Ok(output));
1302 }
1303 expiry.as_mut().poll(context).map(|()| Err(Expired))
1304 })
1305 .await
1306}
1307
1308fn remaining(deadline: Instant) -> Duration {
1309 deadline.saturating_duration_since(Instant::now())
1310}
1311
1312fn deadline_error(operation: &str) -> io::Error {
1313 io::Error::new(
1314 io::ErrorKind::TimedOut,
1315 format!("{operation} deadline expired"),
1316 )
1317}
1318
1319fn report_failure(context: &str, detail: &dyn fmt::Display) {
1324 eprintln!("blazingly-native: {context}: {detail}");
1325}
1326
1327fn resolve_address(address: impl ToSocketAddrs) -> io::Result<SocketAddr> {
1328 address.to_socket_addrs()?.next().ok_or_else(|| {
1329 io::Error::new(
1330 io::ErrorKind::InvalidInput,
1331 "address resolved to no sockets",
1332 )
1333 })
1334}
1335
1336#[derive(Clone, Copy, Eq, PartialEq)]
1338enum Expectation {
1339 None,
1341 Continue,
1343 Unsupported,
1345}
1346
1347fn request_expectation(parsed: &ParsedHead, buffer: &[u8]) -> Expectation {
1353 let mut expectation = Expectation::None;
1354 for header in parsed.headers.iter() {
1355 if !header
1356 .name
1357 .text(buffer)
1358 .is_some_and(|name| name.eq_ignore_ascii_case("expect"))
1359 {
1360 continue;
1361 }
1362 let Some(value) = header.value.text(buffer) else {
1363 return Expectation::Unsupported;
1364 };
1365 for token in value
1366 .split(',')
1367 .map(str::trim)
1368 .filter(|token| !token.is_empty())
1369 {
1370 if !token.eq_ignore_ascii_case("100-continue") {
1371 return Expectation::Unsupported;
1372 }
1373 expectation = Expectation::Continue;
1374 }
1375 }
1376 if expectation == Expectation::Continue && is_http_1_0(parsed, buffer) {
1377 return Expectation::None;
1378 }
1379 expectation
1380}
1381
1382fn is_http_1_0(parsed: &ParsedHead, buffer: &[u8]) -> bool {
1387 buffer
1388 .get(..parsed.head_bytes)
1389 .and_then(|head| head.split(|byte| *byte == b'\n').next())
1390 .is_some_and(|line| line.trim_ascii_end().ends_with(b"HTTP/1.0"))
1391}
1392
1393const fn expects_request_body(body: BodyFraming) -> bool {
1395 match body {
1396 BodyFraming::ContentLength(length) => length > 0,
1397 BodyFraming::Chunked => true,
1398 }
1399}
1400
1401const fn body_exceeds_limit(body: BodyFraming, limits: ServerLimits) -> bool {
1403 match body {
1404 BodyFraming::ContentLength(length) => length > limits.max_body_bytes,
1405 BodyFraming::Chunked => false,
1406 }
1407}
1408
1409async fn write_continue<IO>(io: &mut IO, write_timeout: Duration) -> io::Result<()>
1411where
1412 IO: AsyncWrite + Unpin,
1413{
1414 write_all_within(io, CONTINUE_RESPONSE, write_timeout).await?;
1415 flush_within(io, write_timeout).await
1416}
1417
1418async fn write_continue_native(
1423 io: &mut TcpStream,
1424 wire: &mut Vec<u8>,
1425 write_timeout: Duration,
1426) -> io::Result<()> {
1427 flush_native_pending(io, wire, write_timeout).await?;
1428 wire.clear();
1429 wire.extend_from_slice(CONTINUE_RESPONSE);
1430 native_write_all(io, wire, write_timeout).await
1431}
1432
1433#[allow(clippy::too_many_lines, clippy::too_many_arguments)]
1434async fn serve_connection<IO>(
1435 app: &HttpApp,
1436 limits: ServerLimits,
1437 io: &mut IO,
1438 peer_addr: Option<SocketAddr>,
1439 scheme: &'static str,
1440 shutdown: Option<&ShutdownState>,
1441 request_timeout: Option<Duration>,
1442 ownership: TransportOwnership,
1443) -> io::Result<ConnectionOutcome>
1444where
1445 IO: AsyncRead + AsyncWrite + Unpin,
1446{
1447 ensure_date_updater();
1448 let mut buffer = Vec::with_capacity(READ_CHUNK_BYTES);
1449 let mut read_chunk = vec![0_u8; READ_CHUNK_BYTES];
1450 let mut wire_response = Vec::with_capacity(READ_CHUNK_BYTES);
1451 let mut completed_requests = 0_usize;
1452
1453 loop {
1454 let mut head_deadline = None;
1455 let parsed = loop {
1456 #[cfg(feature = "http2")]
1457 {
1458 if buffer.starts_with(shiguredo_http2::CONNECTION_PREFACE) {
1459 Box::pin(http2::serve_connection(
1460 app,
1461 limits,
1462 io,
1463 peer_addr,
1464 scheme,
1465 shutdown,
1466 std::mem::take(&mut buffer),
1467 request_timeout,
1468 ))
1469 .await?;
1470 return Ok(ConnectionOutcome::Completed);
1471 }
1472 }
1473
1474 if !is_partial_http2_preface(&buffer) {
1475 match parse_head(&buffer, limits) {
1476 Ok(Some(parsed)) => break parsed,
1477 Ok(None) if buffer.len() >= limits.max_header_bytes => {
1478 write_rejection(
1479 io,
1480 &mut wire_response,
1481 431,
1482 "request_header_too_large",
1483 "request headers exceed the configured limit",
1484 limits.write_timeout,
1485 )
1486 .await?;
1487 return Ok(ConnectionOutcome::Completed);
1488 }
1489 Ok(None) => {}
1490 Err(rejection) => {
1491 write_rejection(
1492 io,
1493 &mut wire_response,
1494 rejection.status,
1495 rejection.code,
1496 rejection.message,
1497 limits.write_timeout,
1498 )
1499 .await?;
1500 return Ok(ConnectionOutcome::Completed);
1501 }
1502 }
1503 }
1504
1505 let receiving_head = !buffer.is_empty();
1506 let wait = if receiving_head {
1507 remaining(
1508 *head_deadline
1509 .get_or_insert_with(|| Instant::now() + limits.header_read_timeout),
1510 )
1511 } else {
1512 limits.idle_timeout
1513 };
1514 let Ok(read) = within(wait, io.read(&mut read_chunk)).await else {
1515 if receiving_head {
1516 let _ = write_rejection(
1517 io,
1518 &mut wire_response,
1519 408,
1520 "request_timeout",
1521 "the request head did not arrive within the configured deadline",
1522 limits.write_timeout,
1523 )
1524 .await;
1525 }
1526 return Ok(ConnectionOutcome::Completed);
1527 };
1528 let read = read?;
1529 if read == 0 {
1530 return Ok(ConnectionOutcome::Completed);
1531 }
1532 buffer.extend_from_slice(&read_chunk[..read]);
1533 };
1534
1535 let expectation = request_expectation(&parsed, &buffer);
1538 if expectation == Expectation::Unsupported {
1539 write_rejection(
1540 io,
1541 &mut wire_response,
1542 417,
1543 "expectation_failed",
1544 "the Expect header requested an unsupported expectation",
1545 limits.write_timeout,
1546 )
1547 .await?;
1548 return Ok(ConnectionOutcome::Completed);
1549 }
1550 if body_exceeds_limit(parsed.body, limits) {
1551 write_rejection(
1552 io,
1553 &mut wire_response,
1554 413,
1555 "payload_too_large",
1556 "request body exceeds the configured limit",
1557 limits.write_timeout,
1558 )
1559 .await?;
1560 return Ok(ConnectionOutcome::Completed);
1561 }
1562 if expectation == Expectation::Continue && expects_request_body(parsed.body) {
1563 write_continue(io, limits.write_timeout).await?;
1564 }
1565
1566 let request_method = framework_method(parsed.method);
1567 let streams_request_body = parsed.target.text(&buffer).is_some_and(|target| {
1568 app.request_body_source(request_method, target)
1569 == Some(blazingly_core::InputSource::Stream)
1570 });
1571 if streams_request_body {
1572 let dispatched = {
1573 let mut reader = CompatChunkReader {
1574 io: &mut *io,
1575 scratch: Vec::new(),
1576 };
1577 dispatch_streaming(
1578 app,
1579 limits,
1580 &mut reader,
1581 &mut buffer,
1582 &parsed,
1583 peer_addr,
1584 scheme,
1585 request_timeout,
1586 )
1587 .await
1588 };
1589 let mut response = match dispatched {
1590 Ok(response) => response,
1591 Err(StreamingRequestError::Io(error)) => return Err(error),
1592 Err(StreamingRequestError::Protocol(rejection)) => {
1593 write_rejection(
1594 io,
1595 &mut wire_response,
1596 rejection.status,
1597 rejection.code,
1598 rejection.message,
1599 limits.write_timeout,
1600 )
1601 .await?;
1602 return Ok(ConnectionOutcome::Completed);
1603 }
1604 Err(StreamingRequestError::Incomplete(message)) => {
1605 write_rejection(
1606 io,
1607 &mut wire_response,
1608 400,
1609 "incomplete_body",
1610 message,
1611 limits.write_timeout,
1612 )
1613 .await?;
1614 return Ok(ConnectionOutcome::Completed);
1615 }
1616 };
1617 completed_requests += 1;
1618 let request_limit_reached = limits
1619 .max_requests_per_connection
1620 .is_some_and(|limit| completed_requests >= limit.get());
1621 let keep_alive = parsed.keep_alive
1622 && !request_limit_reached
1623 && !shutdown.is_some_and(|shutdown| shutdown.requested.load(Ordering::Acquire));
1624 let send_body = parsed.method != blazingly_wire::Method::Head
1625 && !matches!(response.status(), 204 | 304)
1626 && !(parsed.method == blazingly_wire::Method::Connect
1627 && (200..300).contains(&response.status()));
1628 let send_content_length = response.status() != 204
1629 && !(parsed.method == blazingly_wire::Method::Connect
1630 && (200..300).contains(&response.status()));
1631 if let Some(upgrade) = response.take_upgrade() {
1632 if ownership == TransportOwnership::Borrowed {
1633 write_rejection(
1634 io,
1635 &mut wire_response,
1636 501,
1637 "upgrade_transport_unsupported",
1638 "this borrowed transport cannot transfer ownership for an upgrade",
1639 limits.write_timeout,
1640 )
1641 .await?;
1642 return Ok(ConnectionOutcome::Completed);
1643 }
1644 wire_response.clear();
1645 with_cached_date(|date| {
1646 blazingly_wire::encode_upgrade_response(
1647 &mut wire_response,
1648 response.headers(),
1649 date,
1650 )
1651 })?;
1652 write_all_within(io, &wire_response, limits.write_timeout).await?;
1653 flush_within(io, limits.write_timeout).await?;
1654 schedule_background(response.take_background_tasks());
1655 return Ok(ConnectionOutcome::Upgraded {
1656 upgrade,
1657 buffered: std::mem::take(&mut buffer),
1658 });
1659 }
1660 write_response(
1661 io,
1662 &mut wire_response,
1663 &mut response,
1664 keep_alive,
1665 send_body,
1666 send_content_length,
1667 limits.write_timeout,
1668 )
1669 .await?;
1670 schedule_background(response.take_background_tasks());
1671 if !keep_alive {
1672 return Ok(ConnectionOutcome::Completed);
1673 }
1674 continue;
1675 }
1676
1677 let mut decoded_chunked = None;
1678 let request_bytes = match parsed.body {
1679 BodyFraming::ContentLength(content_length) => {
1680 if content_length > limits.max_body_bytes {
1681 write_rejection(
1682 io,
1683 &mut wire_response,
1684 413,
1685 "payload_too_large",
1686 "request body exceeds the configured limit",
1687 limits.write_timeout,
1688 )
1689 .await?;
1690 return Ok(ConnectionOutcome::Completed);
1691 }
1692 let request_bytes =
1693 parsed
1694 .head_bytes
1695 .checked_add(content_length)
1696 .ok_or_else(|| {
1697 io::Error::new(io::ErrorKind::InvalidData, "request size overflow")
1698 })?;
1699 while buffer.len() < request_bytes {
1700 let Ok(read) = within(limits.body_read_timeout, io.read(&mut read_chunk)).await
1701 else {
1702 return Err(deadline_error("request body read"));
1703 };
1704 let read = read?;
1705 if read == 0 {
1706 write_rejection(
1707 io,
1708 &mut wire_response,
1709 400,
1710 "incomplete_body",
1711 "request body ended before Content-Length bytes arrived",
1712 limits.write_timeout,
1713 )
1714 .await?;
1715 return Ok(ConnectionOutcome::Completed);
1716 }
1717 buffer.extend_from_slice(&read_chunk[..read]);
1718 }
1719 request_bytes
1720 }
1721 BodyFraming::Chunked => {
1722 let mut decoder = ChunkDecoder::new(parsed.head_bytes, wire_limits(limits));
1723 loop {
1724 match decoder.advance(&buffer) {
1725 Ok(Some(decoded_body)) => {
1726 let consumed = decoded_body.consumed;
1727 decoded_chunked = Some(decoded_body.body);
1728 break consumed;
1729 }
1730 Ok(None) => {}
1731 Err(rejection) => {
1732 write_rejection(
1733 io,
1734 &mut wire_response,
1735 rejection.status,
1736 rejection.code,
1737 rejection.message,
1738 limits.write_timeout,
1739 )
1740 .await?;
1741 return Ok(ConnectionOutcome::Completed);
1742 }
1743 }
1744 let Ok(read) = within(limits.body_read_timeout, io.read(&mut read_chunk)).await
1745 else {
1746 return Err(deadline_error("request body read"));
1747 };
1748 let read = read?;
1749 if read == 0 {
1750 write_rejection(
1751 io,
1752 &mut wire_response,
1753 400,
1754 "incomplete_body",
1755 "chunked request body ended before its final chunk",
1756 limits.write_timeout,
1757 )
1758 .await?;
1759 return Ok(ConnectionOutcome::Completed);
1760 }
1761 buffer.extend_from_slice(&read_chunk[..read]);
1762 }
1763 }
1764 };
1765
1766 let body = decoded_chunked
1767 .as_deref()
1768 .unwrap_or(&buffer[parsed.head_bytes..request_bytes]);
1769 let target = std::str::from_utf8(&buffer[parsed.target.start..parsed.target.end])
1770 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
1771 let native_request = NativeRequest {
1772 method: framework_method(parsed.method),
1773 target,
1774 buffer: &buffer,
1775 headers: &parsed.headers,
1776 body,
1777 peer_addr,
1778 scheme,
1779 };
1780 let mut response = if let Some(timeout) = request_timeout {
1781 app.call_view_controlled(
1782 &native_request,
1783 InvocationControl::new().with_timeout(compio::time::sleep(timeout)),
1784 )
1785 .await
1786 } else {
1787 app.call_view(&native_request).await
1788 };
1789 completed_requests += 1;
1790 let request_limit_reached = limits
1791 .max_requests_per_connection
1792 .is_some_and(|limit| completed_requests >= limit.get());
1793 let keep_alive = parsed.keep_alive
1794 && !request_limit_reached
1795 && !shutdown.is_some_and(|shutdown| shutdown.requested.load(Ordering::Acquire));
1796 let send_body = parsed.method != blazingly_wire::Method::Head
1797 && !matches!(response.status(), 204 | 304)
1798 && !(parsed.method == blazingly_wire::Method::Connect
1799 && (200..300).contains(&response.status()));
1800 let send_content_length = response.status() != 204
1801 && !(parsed.method == blazingly_wire::Method::Connect
1802 && (200..300).contains(&response.status()));
1803 if let Some(upgrade) = response.take_upgrade() {
1804 if ownership == TransportOwnership::Borrowed {
1805 write_rejection(
1806 io,
1807 &mut wire_response,
1808 501,
1809 "upgrade_transport_unsupported",
1810 "this borrowed transport cannot transfer ownership for an upgrade",
1811 limits.write_timeout,
1812 )
1813 .await?;
1814 return Ok(ConnectionOutcome::Completed);
1815 }
1816 wire_response.clear();
1817 with_cached_date(|date| {
1818 blazingly_wire::encode_upgrade_response(
1819 &mut wire_response,
1820 response.headers(),
1821 date,
1822 )
1823 })?;
1824 write_all_within(io, &wire_response, limits.write_timeout).await?;
1825 flush_within(io, limits.write_timeout).await?;
1826 schedule_background(response.take_background_tasks());
1827 let buffered = buffer
1828 .get(request_bytes..)
1829 .map_or_else(Vec::new, <[u8]>::to_vec);
1830 return Ok(ConnectionOutcome::Upgraded { upgrade, buffered });
1831 }
1832 write_response(
1833 io,
1834 &mut wire_response,
1835 &mut response,
1836 keep_alive,
1837 send_body,
1838 send_content_length,
1839 limits.write_timeout,
1840 )
1841 .await?;
1842 schedule_background(response.take_background_tasks());
1843 consume_prefix(&mut buffer, request_bytes);
1844
1845 if !keep_alive {
1846 return Ok(ConnectionOutcome::Completed);
1847 }
1848 }
1849}
1850
1851struct CompatUpgradedIo<IO> {
1857 io: IO,
1858 buffered: Vec<u8>,
1859}
1860
1861impl<IO> UpgradedIo for CompatUpgradedIo<IO>
1862where
1863 IO: AsyncRead + AsyncWrite + Unpin + 'static,
1864{
1865 fn read(
1866 &mut self,
1867 ) -> Pin<Box<dyn Future<Output = Result<Option<Vec<u8>>, UpgradeIoError>> + '_>> {
1868 Box::pin(async move {
1869 if !self.buffered.is_empty() {
1870 return Ok(Some(std::mem::take(&mut self.buffered)));
1871 }
1872 let mut chunk = vec![0_u8; READ_CHUNK_BYTES];
1873 let read = self
1874 .io
1875 .read(&mut chunk)
1876 .await
1877 .map_err(|error| upgrade_io_error("upgrade_read_failed", &error))?;
1878 if read == 0 {
1879 return Ok(None);
1880 }
1881 chunk.truncate(read);
1882 Ok(Some(chunk))
1883 })
1884 }
1885
1886 fn write(
1887 &mut self,
1888 bytes: Vec<u8>,
1889 ) -> Pin<Box<dyn Future<Output = Result<(), UpgradeIoError>> + '_>> {
1890 Box::pin(async move {
1891 self.io
1892 .write_all(&bytes)
1893 .await
1894 .map_err(|error| upgrade_io_error("upgrade_write_failed", &error))?;
1895 self.io
1898 .flush()
1899 .await
1900 .map_err(|error| upgrade_io_error("upgrade_write_failed", &error))
1901 })
1902 }
1903
1904 fn shutdown(&mut self) -> Pin<Box<dyn Future<Output = Result<(), UpgradeIoError>> + '_>> {
1905 Box::pin(async move {
1906 self.io
1907 .close()
1908 .await
1909 .map_err(|error| upgrade_io_error("upgrade_shutdown_failed", &error))
1910 })
1911 }
1912}
1913
1914enum StreamingRequestError {
1916 Io(io::Error),
1917 Protocol(Rejection),
1918 Incomplete(&'static str),
1919}
1920
1921impl From<io::Error> for StreamingRequestError {
1922 fn from(error: io::Error) -> Self {
1923 Self::Io(error)
1924 }
1925}
1926
1927impl From<Rejection> for StreamingRequestError {
1928 fn from(error: Rejection) -> Self {
1929 Self::Protocol(error)
1930 }
1931}
1932
1933async fn send_streaming_chunk(
1934 sender: &IncomingBodySender,
1935 chunk: Vec<u8>,
1936 mut response_future: Pin<&mut dyn Future<Output = Response>>,
1937 response: &mut Option<Response>,
1938) -> bool {
1939 enum Activity {
1940 Sent(bool),
1941 Response(Box<Response>),
1942 }
1943
1944 if response.is_some() {
1945 return false;
1946 }
1947 let mut delivery = std::pin::pin!(sender.send(chunk));
1948 let activity = std::future::poll_fn(|context| {
1949 if let Poll::Ready(response) = response_future.as_mut().poll(context) {
1950 return Poll::Ready(Activity::Response(Box::new(response)));
1951 }
1952 delivery.as_mut().poll(context).map(Activity::Sent)
1953 })
1954 .await;
1955 match activity {
1956 Activity::Sent(sent) => sent,
1957 Activity::Response(completed) => {
1958 *response = Some(*completed);
1959 false
1960 }
1961 }
1962}
1963
1964trait BodyChunkReader {
1977 async fn read_chunk(&mut self, buffer: Vec<u8>, window: usize) -> io::Result<Vec<u8>>;
1980}
1981
1982struct NativeChunkReader<'io> {
1984 io: &'io mut TcpStream,
1985}
1986
1987impl BodyChunkReader for NativeChunkReader<'_> {
1988 async fn read_chunk(&mut self, mut buffer: Vec<u8>, window: usize) -> io::Result<Vec<u8>> {
1989 buffer.clear();
1990 buffer.reserve(window);
1991 let result = CompioAsyncReadExt::append(&mut *self.io, buffer).await;
1992 result.0.map(|_| result.1)
1993 }
1994}
1995
1996struct CompatChunkReader<'io, IO> {
2003 io: &'io mut IO,
2004 scratch: Vec<u8>,
2005}
2006
2007impl<IO> BodyChunkReader for CompatChunkReader<'_, IO>
2008where
2009 IO: AsyncRead + Unpin,
2010{
2011 async fn read_chunk(&mut self, mut buffer: Vec<u8>, window: usize) -> io::Result<Vec<u8>> {
2012 buffer.clear();
2013 if self.scratch.len() < window {
2014 self.scratch.resize(window, 0);
2015 }
2016 let read = self.io.read(&mut self.scratch[..window]).await?;
2017 buffer.extend_from_slice(&self.scratch[..read]);
2018 Ok(buffer)
2019 }
2020}
2021
2022async fn read_body_chunk<R>(
2023 reader: &mut R,
2024 buffer: Vec<u8>,
2025 window: usize,
2026 mut response_future: Pin<&mut dyn Future<Output = Response>>,
2027 response: &mut Option<Response>,
2028) -> io::Result<Vec<u8>>
2029where
2030 R: BodyChunkReader,
2031{
2032 let mut read = std::pin::pin!(reader.read_chunk(buffer, window));
2033 if response.is_none() {
2034 enum Activity {
2035 Read(io::Result<Vec<u8>>),
2036 Response(Box<Response>),
2037 }
2038 let activity = std::future::poll_fn(|context| {
2039 if let Poll::Ready(response) = response_future.as_mut().poll(context) {
2040 return Poll::Ready(Activity::Response(Box::new(response)));
2041 }
2042 read.as_mut().poll(context).map(Activity::Read)
2043 })
2044 .await;
2045 match activity {
2046 Activity::Read(result) => return result,
2047 Activity::Response(completed) => *response = Some(*completed),
2048 }
2049 }
2050 read.await
2051}
2052
2053#[allow(clippy::too_many_arguments, clippy::too_many_lines)]
2054async fn dispatch_streaming<R>(
2055 app: &HttpApp,
2056 limits: ServerLimits,
2057 reader: &mut R,
2058 buffer: &mut Vec<u8>,
2059 parsed: &ParsedHead,
2060 peer_addr: Option<SocketAddr>,
2061 scheme: &'static str,
2062 request_timeout: Option<Duration>,
2063) -> Result<Response, StreamingRequestError>
2064where
2065 R: BodyChunkReader,
2066{
2067 let window = limits.stream_read_bytes;
2068 let exact_length = match parsed.body {
2069 BodyFraming::ContentLength(length) => {
2070 if length > limits.max_body_bytes {
2071 return Err(StreamingRequestError::Protocol(Rejection {
2072 status: 413,
2073 code: "payload_too_large",
2074 message: "request body exceeds the configured limit",
2075 }));
2076 }
2077 Some(u64::try_from(length).unwrap_or(u64::MAX))
2078 }
2079 BodyFraming::Chunked => None,
2080 };
2081 let (sender, body) = incoming_body_channel(exact_length, window.saturating_mul(2));
2082 let request = owned_request(parsed, buffer, body, peer_addr, scheme)?;
2083 consume_prefix(buffer, parsed.head_bytes);
2084 let mut response_future: Pin<Box<dyn Future<Output = Response>>> =
2085 if let Some(timeout) = request_timeout {
2086 Box::pin(app.call_view_controlled(
2087 &request,
2088 InvocationControl::new().with_timeout(compio::time::sleep(timeout)),
2089 ))
2090 } else {
2091 Box::pin(app.call_view(&request))
2092 };
2093 let mut response = None;
2094 let mut receiver_open = true;
2095
2096 match parsed.body {
2097 BodyFraming::ContentLength(length) => {
2098 let mut remaining = length;
2099 while remaining > 0 && !buffer.is_empty() {
2102 let available = remaining.min(buffer.len()).min(window);
2103 let mut chunk = sender.take_spare();
2104 chunk.extend_from_slice(&buffer[..available]);
2105 consume_prefix(buffer, available);
2106 remaining -= available;
2107 if receiver_open {
2108 receiver_open = send_streaming_chunk(
2109 &sender,
2110 chunk,
2111 response_future.as_mut(),
2112 &mut response,
2113 )
2114 .await;
2115 }
2116 }
2117 while remaining > 0 {
2120 let spare = sender.take_spare();
2121 let Ok(chunk) = within(
2122 limits.body_read_timeout,
2123 read_body_chunk(
2124 &mut *reader,
2125 spare,
2126 window,
2127 response_future.as_mut(),
2128 &mut response,
2129 ),
2130 )
2131 .await
2132 else {
2133 sender.fail(BodyStreamError::new(
2134 "upload_timeout",
2135 "request body stalled past the configured deadline",
2136 ));
2137 return Err(StreamingRequestError::Io(deadline_error(
2138 "request body read",
2139 )));
2140 };
2141 let mut chunk = chunk?;
2142 if chunk.is_empty() {
2143 sender.fail(BodyStreamError::new(
2144 "incomplete_upload",
2145 "request body ended before Content-Length bytes arrived",
2146 ));
2147 return Err(StreamingRequestError::Incomplete(
2148 "request body ended before Content-Length bytes arrived",
2149 ));
2150 }
2151 if chunk.len() > remaining {
2152 buffer.extend_from_slice(&chunk[remaining..]);
2155 chunk.truncate(remaining);
2156 }
2157 remaining -= chunk.len();
2158 if receiver_open {
2159 receiver_open = send_streaming_chunk(
2160 &sender,
2161 chunk,
2162 response_future.as_mut(),
2163 &mut response,
2164 )
2165 .await;
2166 }
2167 }
2168 }
2169 BodyFraming::Chunked => {
2170 let mut decoder = StreamingChunkDecoder::new(0, wire_limits(limits));
2171 loop {
2172 match decoder.advance(buffer)? {
2173 StreamingChunk::Data(range) => {
2174 let mut chunk = sender.take_spare();
2175 chunk.extend_from_slice(
2176 range
2177 .bytes(buffer)
2178 .expect("decoder ranges stay inside the receive buffer"),
2179 );
2180 let consumed = decoder.consumed_prefix();
2181 consume_prefix(buffer, consumed);
2182 decoder.discard_prefix(consumed);
2183 if receiver_open {
2184 receiver_open = send_streaming_chunk(
2185 &sender,
2186 chunk,
2187 response_future.as_mut(),
2188 &mut response,
2189 )
2190 .await;
2191 }
2192 }
2193 StreamingChunk::Complete { consumed } => {
2194 consume_prefix(buffer, consumed);
2195 break;
2196 }
2197 StreamingChunk::NeedMore => {
2198 let spare = sender.take_spare();
2199 let Ok(chunk) = within(
2200 limits.body_read_timeout,
2201 read_body_chunk(
2202 &mut *reader,
2203 spare,
2204 window,
2205 response_future.as_mut(),
2206 &mut response,
2207 ),
2208 )
2209 .await
2210 else {
2211 sender.fail(BodyStreamError::new(
2212 "upload_timeout",
2213 "request body stalled past the configured deadline",
2214 ));
2215 return Err(StreamingRequestError::Io(deadline_error(
2216 "request body read",
2217 )));
2218 };
2219 let chunk = chunk?;
2220 if chunk.is_empty() {
2221 sender.fail(BodyStreamError::new(
2222 "incomplete_upload",
2223 "chunked request body ended before its final chunk",
2224 ));
2225 return Err(StreamingRequestError::Incomplete(
2226 "chunked request body ended before its final chunk",
2227 ));
2228 }
2229 buffer.extend_from_slice(&chunk);
2230 sender.store_spare(chunk);
2231 }
2232 }
2233 }
2234 }
2235 }
2236 sender.close();
2237 Ok(match response {
2238 Some(response) => response,
2239 None => response_future.await,
2240 })
2241}
2242
2243#[allow(clippy::too_many_lines)]
2250async fn serve_native_connection(
2251 app: &HttpApp,
2252 limits: ServerLimits,
2253 mut io: TcpStream,
2254 peer_addr: Option<SocketAddr>,
2255 scheme: &'static str,
2256 shutdown: Option<&ShutdownState>,
2257 request_timeout: Option<Duration>,
2258) -> io::Result<()> {
2259 ensure_date_updater();
2260 let mut buffer = Vec::with_capacity(READ_CHUNK_BYTES);
2261 let mut wire_response = Vec::with_capacity(READ_CHUNK_BYTES);
2262 let mut completed_requests = 0_usize;
2263 let mut buffered_responses = 0_usize;
2264
2265 loop {
2266 let mut head_deadline = None;
2267 let parsed = loop {
2268 let connection_start = completed_requests == 0;
2272 #[cfg(feature = "http2")]
2273 if connection_start && buffer.starts_with(shiguredo_http2::CONNECTION_PREFACE) {
2274 let initial = std::mem::take(&mut buffer);
2275 let mut stream = Box::pin(AsyncStream::new(io));
2276 return Box::pin(http2::serve_connection(
2277 app,
2278 limits,
2279 &mut stream,
2280 peer_addr,
2281 scheme,
2282 shutdown,
2283 initial,
2284 request_timeout,
2285 ))
2286 .await;
2287 }
2288
2289 if !connection_start || !is_partial_http2_preface(&buffer) {
2290 match parse_head(&buffer, limits) {
2291 Ok(Some(parsed)) => break parsed,
2292 Ok(None) if buffer.len() >= limits.max_header_bytes => {
2293 write_rejection_native(
2294 &mut io,
2295 &mut wire_response,
2296 431,
2297 "request_header_too_large",
2298 "request headers exceed the configured limit",
2299 limits.write_timeout,
2300 )
2301 .await?;
2302 return Ok(());
2303 }
2304 Ok(None) => {}
2305 Err(rejection) => {
2306 write_rejection_native(
2307 &mut io,
2308 &mut wire_response,
2309 rejection.status,
2310 rejection.code,
2311 rejection.message,
2312 limits.write_timeout,
2313 )
2314 .await?;
2315 return Ok(());
2316 }
2317 }
2318 }
2319
2320 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2321 buffered_responses = 0;
2322 let receiving_head = !buffer.is_empty();
2323 let wait = if receiving_head {
2324 remaining(
2325 *head_deadline
2326 .get_or_insert_with(|| Instant::now() + limits.header_read_timeout),
2327 )
2328 } else {
2329 limits.idle_timeout
2330 };
2331 let Ok(read) = within(wait, native_read_more(&mut io, &mut buffer)).await else {
2332 if receiving_head {
2333 let _ = write_rejection_native(
2334 &mut io,
2335 &mut wire_response,
2336 408,
2337 "request_timeout",
2338 "the request head did not arrive within the configured deadline",
2339 limits.write_timeout,
2340 )
2341 .await;
2342 }
2343 return Ok(());
2344 };
2345 if read? == 0 {
2346 return Ok(());
2347 }
2348 };
2349
2350 let expectation = request_expectation(&parsed, &buffer);
2353 if expectation == Expectation::Unsupported {
2354 write_rejection_native(
2355 &mut io,
2356 &mut wire_response,
2357 417,
2358 "expectation_failed",
2359 "the Expect header requested an unsupported expectation",
2360 limits.write_timeout,
2361 )
2362 .await?;
2363 return Ok(());
2364 }
2365 if body_exceeds_limit(parsed.body, limits) {
2366 write_rejection_native(
2367 &mut io,
2368 &mut wire_response,
2369 413,
2370 "payload_too_large",
2371 "request body exceeds the configured limit",
2372 limits.write_timeout,
2373 )
2374 .await?;
2375 return Ok(());
2376 }
2377 if expectation == Expectation::Continue && expects_request_body(parsed.body) {
2378 write_continue_native(&mut io, &mut wire_response, limits.write_timeout).await?;
2379 buffered_responses = 0;
2380 }
2381
2382 let request_method = framework_method(parsed.method);
2383 let request_target = parsed
2384 .target
2385 .text(&buffer)
2386 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid request target"))?;
2387 if app.request_body_source(request_method, request_target)
2388 == Some(blazingly_core::InputSource::Stream)
2389 {
2390 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2391 let dispatched = {
2392 let mut reader = NativeChunkReader { io: &mut io };
2393 dispatch_streaming(
2394 app,
2395 limits,
2396 &mut reader,
2397 &mut buffer,
2398 &parsed,
2399 peer_addr,
2400 scheme,
2401 request_timeout,
2402 )
2403 .await
2404 };
2405 let mut response = match dispatched {
2406 Ok(response) => response,
2407 Err(StreamingRequestError::Io(error)) => return Err(error),
2408 Err(StreamingRequestError::Protocol(rejection)) => {
2409 write_rejection_native(
2410 &mut io,
2411 &mut wire_response,
2412 rejection.status,
2413 rejection.code,
2414 rejection.message,
2415 limits.write_timeout,
2416 )
2417 .await?;
2418 return Ok(());
2419 }
2420 Err(StreamingRequestError::Incomplete(message)) => {
2421 write_rejection_native(
2422 &mut io,
2423 &mut wire_response,
2424 400,
2425 "incomplete_body",
2426 message,
2427 limits.write_timeout,
2428 )
2429 .await?;
2430 return Ok(());
2431 }
2432 };
2433 if let Some(upgrade) = response.take_upgrade() {
2434 with_cached_date(|date| {
2435 blazingly_wire::encode_upgrade_response(
2436 &mut wire_response,
2437 response.headers(),
2438 date,
2439 )
2440 })?;
2441 native_write_all(&mut io, &mut wire_response, limits.write_timeout).await?;
2442 schedule_background(response.take_background_tasks());
2443 return upgrade
2444 .run(Box::new(NativeUpgradedIo {
2445 io,
2446 buffered: buffer,
2447 }))
2448 .await
2449 .map_err(|error| io::Error::other(error.to_string()));
2450 }
2451 completed_requests += 1;
2452 let request_limit_reached = limits
2453 .max_requests_per_connection
2454 .is_some_and(|limit| completed_requests >= limit.get());
2455 let keep_alive = parsed.keep_alive
2456 && !request_limit_reached
2457 && !shutdown.is_some_and(|shutdown| shutdown.requested.load(Ordering::Acquire));
2458 let send_body = parsed.method != blazingly_wire::Method::Head
2459 && !matches!(response.status(), 204 | 304)
2460 && !(parsed.method == blazingly_wire::Method::Connect
2461 && (200..300).contains(&response.status()));
2462 let send_content_length = response.status() != 204
2463 && !(parsed.method == blazingly_wire::Method::Connect
2464 && (200..300).contains(&response.status()));
2465 let streaming_response = response.is_streaming();
2466 write_response_native(
2467 &mut io,
2468 &mut wire_response,
2469 &mut response,
2470 keep_alive,
2471 send_body,
2472 send_content_length,
2473 limits.write_timeout,
2474 )
2475 .await?;
2476 schedule_background(response.take_background_tasks());
2477 if streaming_response {
2478 buffered_responses = 0;
2479 } else {
2480 buffered_responses = 1;
2481 }
2482 if !keep_alive {
2483 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2484 return Ok(());
2485 }
2486 continue;
2487 }
2488
2489 let mut decoded_chunked = None;
2490 let request_bytes = match parsed.body {
2491 BodyFraming::ContentLength(content_length) => {
2492 if content_length > limits.max_body_bytes {
2493 write_rejection_native(
2494 &mut io,
2495 &mut wire_response,
2496 413,
2497 "payload_too_large",
2498 "request body exceeds the configured limit",
2499 limits.write_timeout,
2500 )
2501 .await?;
2502 return Ok(());
2503 }
2504 let request_bytes =
2505 parsed
2506 .head_bytes
2507 .checked_add(content_length)
2508 .ok_or_else(|| {
2509 io::Error::new(io::ErrorKind::InvalidData, "request size overflow")
2510 })?;
2511 while buffer.len() < request_bytes {
2512 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2513 buffered_responses = 0;
2514 let Ok(read) = within(
2515 limits.body_read_timeout,
2516 native_read_more(&mut io, &mut buffer),
2517 )
2518 .await
2519 else {
2520 return Err(deadline_error("request body read"));
2521 };
2522 if read? == 0 {
2523 write_rejection_native(
2524 &mut io,
2525 &mut wire_response,
2526 400,
2527 "incomplete_body",
2528 "request body ended before Content-Length bytes arrived",
2529 limits.write_timeout,
2530 )
2531 .await?;
2532 return Ok(());
2533 }
2534 }
2535 request_bytes
2536 }
2537 BodyFraming::Chunked => {
2538 let mut decoder = ChunkDecoder::new(parsed.head_bytes, wire_limits(limits));
2539 loop {
2540 match decoder.advance(&buffer) {
2541 Ok(Some(decoded_body)) => {
2542 let consumed = decoded_body.consumed;
2543 decoded_chunked = Some(decoded_body.body);
2544 break consumed;
2545 }
2546 Ok(None) => {}
2547 Err(rejection) => {
2548 write_rejection_native(
2549 &mut io,
2550 &mut wire_response,
2551 rejection.status,
2552 rejection.code,
2553 rejection.message,
2554 limits.write_timeout,
2555 )
2556 .await?;
2557 return Ok(());
2558 }
2559 }
2560 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2561 buffered_responses = 0;
2562 let Ok(read) = within(
2563 limits.body_read_timeout,
2564 native_read_more(&mut io, &mut buffer),
2565 )
2566 .await
2567 else {
2568 return Err(deadline_error("request body read"));
2569 };
2570 if read? == 0 {
2571 write_rejection_native(
2572 &mut io,
2573 &mut wire_response,
2574 400,
2575 "incomplete_body",
2576 "chunked request body ended before its final chunk",
2577 limits.write_timeout,
2578 )
2579 .await?;
2580 return Ok(());
2581 }
2582 }
2583 }
2584 };
2585
2586 let body = decoded_chunked
2587 .as_deref()
2588 .unwrap_or(&buffer[parsed.head_bytes..request_bytes]);
2589 let target = std::str::from_utf8(&buffer[parsed.target.start..parsed.target.end])
2590 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
2591 let native_request = NativeRequest {
2592 method: framework_method(parsed.method),
2593 target,
2594 buffer: &buffer,
2595 headers: &parsed.headers,
2596 body,
2597 peer_addr,
2598 scheme,
2599 };
2600 let mut response = if let Some(timeout) = request_timeout {
2601 app.call_view_controlled(
2602 &native_request,
2603 InvocationControl::new().with_timeout(compio::time::sleep(timeout)),
2604 )
2605 .await
2606 } else {
2607 app.call_view(&native_request).await
2608 };
2609 if let Some(upgrade) = response.take_upgrade() {
2610 with_cached_date(|date| {
2611 blazingly_wire::encode_upgrade_response(
2612 &mut wire_response,
2613 response.headers(),
2614 date,
2615 )
2616 })?;
2617 native_write_all(&mut io, &mut wire_response, limits.write_timeout).await?;
2618 schedule_background(response.take_background_tasks());
2619 let buffered = buffer
2620 .get(request_bytes..)
2621 .map_or_else(Vec::new, <[u8]>::to_vec);
2622 return upgrade
2623 .run(Box::new(NativeUpgradedIo { io, buffered }))
2624 .await
2625 .map_err(|error| io::Error::other(error.to_string()));
2626 }
2627 completed_requests += 1;
2628 let request_limit_reached = limits
2629 .max_requests_per_connection
2630 .is_some_and(|limit| completed_requests >= limit.get());
2631 let keep_alive = parsed.keep_alive
2632 && !request_limit_reached
2633 && !shutdown.is_some_and(|shutdown| shutdown.requested.load(Ordering::Acquire));
2634 let send_body = parsed.method != blazingly_wire::Method::Head
2635 && !matches!(response.status(), 204 | 304)
2636 && !(parsed.method == blazingly_wire::Method::Connect
2637 && (200..300).contains(&response.status()));
2638 let send_content_length = response.status() != 204
2639 && !(parsed.method == blazingly_wire::Method::Connect
2640 && (200..300).contains(&response.status()));
2641 let streaming = response.is_streaming();
2642 write_response_native(
2643 &mut io,
2644 &mut wire_response,
2645 &mut response,
2646 keep_alive,
2647 send_body,
2648 send_content_length,
2649 limits.write_timeout,
2650 )
2651 .await?;
2652 schedule_background(response.take_background_tasks());
2653 if streaming {
2654 buffered_responses = 0;
2655 } else {
2656 buffered_responses += 1;
2657 if buffered_responses >= limits.max_pipeline_batch
2658 || wire_response.len() >= MAX_PIPELINE_WRITE_BYTES
2659 {
2660 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2661 buffered_responses = 0;
2662 }
2663 }
2664 consume_prefix(&mut buffer, request_bytes);
2665
2666 if !keep_alive {
2667 flush_native_pending(&mut io, &mut wire_response, limits.write_timeout).await?;
2668 return Ok(());
2669 }
2670 }
2671}
2672
2673async fn native_read_more(io: &mut TcpStream, buffer: &mut Vec<u8>) -> io::Result<usize> {
2674 let mut taken = std::mem::take(buffer);
2675 if taken.capacity() - taken.len() < READ_CHUNK_BYTES {
2681 taken.reserve(READ_CHUNK_BYTES);
2682 }
2683 let result = CompioAsyncReadExt::append(io, taken).await;
2684 *buffer = result.1;
2685 result.0
2686}
2687
2688struct NativeUpgradedIo {
2689 io: TcpStream,
2690 buffered: Vec<u8>,
2691}
2692
2693impl UpgradedIo for NativeUpgradedIo {
2694 fn read(
2695 &mut self,
2696 ) -> Pin<Box<dyn Future<Output = Result<Option<Vec<u8>>, UpgradeIoError>> + '_>> {
2697 Box::pin(async move {
2698 if !self.buffered.is_empty() {
2699 return Ok(Some(std::mem::take(&mut self.buffered)));
2700 }
2701 let result =
2702 CompioAsyncReadExt::append(&mut self.io, Vec::with_capacity(READ_CHUNK_BYTES))
2703 .await;
2704 let read = result
2705 .0
2706 .map_err(|error| upgrade_io_error("upgrade_read_failed", &error))?;
2707 if read == 0 {
2708 return Ok(None);
2709 }
2710 Ok(Some(result.1))
2711 })
2712 }
2713
2714 fn write(
2715 &mut self,
2716 bytes: Vec<u8>,
2717 ) -> Pin<Box<dyn Future<Output = Result<(), UpgradeIoError>> + '_>> {
2718 Box::pin(async move {
2719 CompioAsyncWriteExt::write_all(&mut self.io, bytes)
2720 .await
2721 .0
2722 .map_err(|error| upgrade_io_error("upgrade_write_failed", &error))
2723 })
2724 }
2725
2726 fn shutdown(&mut self) -> Pin<Box<dyn Future<Output = Result<(), UpgradeIoError>> + '_>> {
2727 Box::pin(async { Ok(()) })
2728 }
2729}
2730
2731fn upgrade_io_error(code: &'static str, error: &io::Error) -> UpgradeIoError {
2732 UpgradeIoError::new(code, error.to_string())
2733}
2734
2735fn schedule_background(tasks: Vec<BackgroundTask>) {
2736 if tasks.is_empty() {
2737 return;
2738 }
2739 let accounting = drain_counter().map(ActiveWork::acquire);
2740 spawn(async move {
2741 for task in tasks {
2742 if let Err(error) = task.run().await {
2743 report_failure("background task failed", &error);
2744 }
2745 }
2746 drop(accounting);
2747 })
2748 .detach();
2749}
2750
2751#[cfg(feature = "http2")]
2752fn is_partial_http2_preface(buffer: &[u8]) -> bool {
2753 !buffer.is_empty()
2754 && buffer.len() < shiguredo_http2::CONNECTION_PREFACE.len()
2755 && shiguredo_http2::CONNECTION_PREFACE.starts_with(buffer)
2756}
2757
2758#[cfg(not(feature = "http2"))]
2759const fn is_partial_http2_preface(_buffer: &[u8]) -> bool {
2760 false
2761}
2762
2763type ParsedHead = blazingly_wire::RequestHead;
2764type Rejection = blazingly_wire::ParseError;
2765
2766fn parse_head(buffer: &[u8], limits: ServerLimits) -> Result<Option<ParsedHead>, Rejection> {
2767 blazingly_wire::parse_request_head(buffer, wire_limits(limits))
2768}
2769
2770const fn wire_limits(limits: ServerLimits) -> blazingly_wire::Limits {
2771 blazingly_wire::Limits::new()
2772 .with_max_header_bytes(limits.max_header_bytes)
2773 .with_max_headers(limits.max_headers)
2774 .with_max_body_bytes(limits.max_body_bytes)
2775 .with_max_chunks(limits.max_chunks)
2776}
2777
2778const fn framework_method(method: blazingly_wire::Method) -> HttpMethod {
2779 match method {
2780 blazingly_wire::Method::Get => HttpMethod::Get,
2781 blazingly_wire::Method::Head => HttpMethod::Head,
2782 blazingly_wire::Method::Post => HttpMethod::Post,
2783 blazingly_wire::Method::Put => HttpMethod::Put,
2784 blazingly_wire::Method::Patch => HttpMethod::Patch,
2785 blazingly_wire::Method::Delete => HttpMethod::Delete,
2786 blazingly_wire::Method::Options => HttpMethod::Options,
2787 blazingly_wire::Method::Trace => HttpMethod::Trace,
2788 blazingly_wire::Method::Connect => HttpMethod::Connect,
2789 }
2790}
2791
2792#[cfg(feature = "http2")]
2793fn parse_method(method: &str) -> Result<HttpMethod, Rejection> {
2794 blazingly_wire::Method::parse(method).map(framework_method)
2795}
2796
2797struct IncomingBodyState {
2798 chunks: VecDeque<Result<Vec<u8>, BodyStreamError>>,
2799 queued_bytes: usize,
2800 #[cfg(feature = "http2")]
2803 consumed_bytes: usize,
2804 max_queued_bytes: usize,
2805 spare: Vec<Vec<u8>>,
2808 closed: bool,
2809 receiver_open: bool,
2810 consumer_waker: Option<Waker>,
2811 producer_waker: Option<Waker>,
2812}
2813
2814impl IncomingBodyState {
2815 fn store_spare(&mut self, mut spent: Vec<u8>) {
2816 if spent.capacity() > 0 && self.spare.len() < MAX_SPARE_BUFFERS {
2817 spent.clear();
2818 self.spare.push(spent);
2819 }
2820 }
2821}
2822
2823#[derive(Clone)]
2824struct IncomingBodySender {
2825 state: Rc<RefCell<IncomingBodyState>>,
2826}
2827
2828impl IncomingBodySender {
2829 async fn send(&self, bytes: Vec<u8>) -> bool {
2830 let mut bytes = Some(bytes);
2831 std::future::poll_fn(|context| {
2832 let mut state = self.state.borrow_mut();
2833 if !state.receiver_open {
2834 return Poll::Ready(false);
2835 }
2836 let length = bytes.as_ref().map_or(0, Vec::len);
2837 if state.chunks.is_empty()
2838 || state.queued_bytes.saturating_add(length) <= state.max_queued_bytes
2839 {
2840 state.queued_bytes = state.queued_bytes.saturating_add(length);
2841 state
2842 .chunks
2843 .push_back(Ok(bytes.take().expect("body chunk is sent once")));
2844 if let Some(waker) = state.consumer_waker.take() {
2845 waker.wake();
2846 }
2847 return Poll::Ready(true);
2848 }
2849 state.producer_waker = Some(context.waker().clone());
2850 Poll::Pending
2851 })
2852 .await
2853 }
2854
2855 #[cfg(feature = "http2")]
2861 fn push(&self, bytes: Vec<u8>) -> bool {
2862 let mut state = self.state.borrow_mut();
2863 if !state.receiver_open {
2864 return false;
2865 }
2866 state.queued_bytes = state.queued_bytes.saturating_add(bytes.len());
2867 state.chunks.push_back(Ok(bytes));
2868 if let Some(waker) = state.consumer_waker.take() {
2869 waker.wake();
2870 }
2871 true
2872 }
2873
2874 #[cfg(feature = "http2")]
2879 async fn consumed(&self) -> (usize, bool) {
2880 std::future::poll_fn(|context| {
2881 let mut state = self.state.borrow_mut();
2882 let consumed = std::mem::take(&mut state.consumed_bytes);
2883 let finished = !state.receiver_open || (state.closed && state.chunks.is_empty());
2884 if consumed > 0 || finished {
2885 return Poll::Ready((consumed, finished));
2886 }
2887 state.producer_waker = Some(context.waker().clone());
2888 Poll::Pending
2889 })
2890 .await
2891 }
2892
2893 fn close(&self) {
2894 let mut state = self.state.borrow_mut();
2895 state.closed = true;
2896 if let Some(waker) = state.consumer_waker.take() {
2897 waker.wake();
2898 }
2899 if let Some(waker) = state.producer_waker.take() {
2900 waker.wake();
2901 }
2902 }
2903
2904 fn take_spare(&self) -> Vec<u8> {
2906 self.state.borrow_mut().spare.pop().unwrap_or_default()
2907 }
2908
2909 fn store_spare(&self, spent: Vec<u8>) {
2911 self.state.borrow_mut().store_spare(spent);
2912 }
2913
2914 fn fail(&self, error: BodyStreamError) {
2915 let mut state = self.state.borrow_mut();
2916 if state.receiver_open {
2917 state.chunks.push_back(Err(error));
2918 }
2919 state.closed = true;
2920 if let Some(waker) = state.consumer_waker.take() {
2921 waker.wake();
2922 }
2923 if let Some(waker) = state.producer_waker.take() {
2924 waker.wake();
2925 }
2926 }
2927}
2928
2929struct NativeIncomingBody {
2930 state: Rc<RefCell<IncomingBodyState>>,
2931}
2932
2933impl BodyStream for NativeIncomingBody {
2934 fn poll_next(
2935 self: Pin<&mut Self>,
2936 context: &mut Context<'_>,
2937 ) -> Poll<Option<Result<Vec<u8>, BodyStreamError>>> {
2938 let mut state = self.state.borrow_mut();
2939 if let Some(chunk) = state.chunks.pop_front() {
2940 if let Ok(bytes) = &chunk {
2941 state.queued_bytes = state.queued_bytes.saturating_sub(bytes.len());
2942 #[cfg(feature = "http2")]
2943 {
2944 state.consumed_bytes = state.consumed_bytes.saturating_add(bytes.len());
2945 }
2946 }
2947 if let Some(waker) = state.producer_waker.take() {
2948 waker.wake();
2949 }
2950 return Poll::Ready(Some(chunk));
2951 }
2952 if state.closed {
2953 return Poll::Ready(None);
2954 }
2955 state.consumer_waker = Some(context.waker().clone());
2956 Poll::Pending
2957 }
2958
2959 fn recycle(self: Pin<&mut Self>, spent: Vec<u8>) {
2960 self.state.borrow_mut().store_spare(spent);
2961 }
2962}
2963
2964impl Drop for NativeIncomingBody {
2965 fn drop(&mut self) {
2966 let mut state = self.state.borrow_mut();
2967 state.receiver_open = false;
2968 if let Some(waker) = state.producer_waker.take() {
2969 waker.wake();
2970 }
2971 }
2972}
2973
2974fn incoming_body_channel(
2975 exact_length: Option<u64>,
2976 max_queued_bytes: usize,
2977) -> (IncomingBodySender, StreamingBody) {
2978 let state = Rc::new(RefCell::new(IncomingBodyState {
2979 chunks: VecDeque::new(),
2980 queued_bytes: 0,
2981 #[cfg(feature = "http2")]
2982 consumed_bytes: 0,
2983 max_queued_bytes,
2984 spare: Vec::new(),
2985 closed: false,
2986 receiver_open: true,
2987 consumer_waker: None,
2988 producer_waker: None,
2989 }));
2990 let sender = IncomingBodySender {
2991 state: Rc::clone(&state),
2992 };
2993 let mut body = StreamingBody::new(NativeIncomingBody { state });
2994 if let Some(exact_length) = exact_length {
2995 body = body.with_exact_length(exact_length);
2996 }
2997 (sender, body)
2998}
2999
3000struct OwnedNativeRequest {
3001 method: HttpMethod,
3002 target: String,
3003 headers: Vec<(String, String)>,
3004 body: RefCell<Option<StreamingBody>>,
3005 peer_addr: Option<SocketAddr>,
3006 scheme: &'static str,
3007}
3008
3009impl HttpRequestView for OwnedNativeRequest {
3010 fn method(&self) -> HttpMethod {
3011 self.method
3012 }
3013
3014 fn target(&self) -> &str {
3015 &self.target
3016 }
3017
3018 fn header_value(&self, name: &str, index: usize) -> Option<&str> {
3019 self.headers
3020 .iter()
3021 .filter(|(header_name, _)| header_name.eq_ignore_ascii_case(name))
3022 .nth(index)
3023 .map(|(_, value)| value.as_str())
3024 }
3025
3026 fn body(&self) -> &[u8] {
3027 &[]
3028 }
3029
3030 fn take_body_stream(&self) -> Option<StreamingBody> {
3031 self.body.borrow_mut().take()
3032 }
3033
3034 fn peer_addr(&self) -> Option<SocketAddr> {
3035 self.peer_addr
3036 }
3037
3038 fn scheme(&self) -> &str {
3039 self.scheme
3040 }
3041}
3042
3043fn owned_request(
3044 parsed: &ParsedHead,
3045 buffer: &[u8],
3046 body: StreamingBody,
3047 peer_addr: Option<SocketAddr>,
3048 scheme: &'static str,
3049) -> io::Result<OwnedNativeRequest> {
3050 let target = parsed
3051 .target
3052 .text(buffer)
3053 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "request target is not UTF-8"))?
3054 .to_owned();
3055 let headers = parsed
3056 .headers
3057 .iter()
3058 .map(|header| {
3059 let name = header.name.text(buffer).ok_or_else(|| {
3060 io::Error::new(
3061 io::ErrorKind::InvalidData,
3062 "request header name is not UTF-8",
3063 )
3064 })?;
3065 let value = header.value.text(buffer).ok_or_else(|| {
3066 io::Error::new(
3067 io::ErrorKind::InvalidData,
3068 "request header value is not UTF-8",
3069 )
3070 })?;
3071 Ok((name.to_owned(), value.to_owned()))
3072 })
3073 .collect::<io::Result<Vec<_>>>()?;
3074 Ok(OwnedNativeRequest {
3075 method: framework_method(parsed.method),
3076 target,
3077 headers,
3078 body: RefCell::new(Some(body)),
3079 peer_addr,
3080 scheme,
3081 })
3082}
3083
3084struct NativeRequest<'request> {
3085 method: HttpMethod,
3086 target: &'request str,
3087 buffer: &'request [u8],
3088 headers: &'request HeaderPositions,
3089 body: &'request [u8],
3090 peer_addr: Option<SocketAddr>,
3091 scheme: &'static str,
3092}
3093
3094impl HttpRequestView for NativeRequest<'_> {
3095 fn method(&self) -> HttpMethod {
3096 self.method
3097 }
3098
3099 fn target(&self) -> &str {
3100 self.target
3101 }
3102
3103 fn header_value(&self, name: &str, index: usize) -> Option<&str> {
3104 self.headers
3105 .iter()
3106 .filter(|header| {
3107 self.buffer
3108 .get(header.name.start..header.name.end)
3109 .and_then(|header| std::str::from_utf8(header).ok())
3110 .is_some_and(|header| native_header_name_matches(header, name))
3111 })
3112 .nth(index)
3113 .and_then(|header| self.buffer.get(header.value.start..header.value.end))
3114 .and_then(|value| std::str::from_utf8(value).ok())
3115 }
3116
3117 fn body(&self) -> &[u8] {
3118 self.body
3119 }
3120
3121 fn peer_addr(&self) -> Option<SocketAddr> {
3122 self.peer_addr
3123 }
3124
3125 fn scheme(&self) -> &str {
3126 self.scheme
3127 }
3128}
3129
3130fn native_header_name_matches(header: &str, argument: &str) -> bool {
3131 header
3132 .bytes()
3133 .map(|byte| byte.to_ascii_lowercase())
3134 .eq(argument
3135 .bytes()
3136 .map(|byte| if byte == b'_' { b'-' } else { byte })
3137 .map(|byte| byte.to_ascii_lowercase()))
3138}
3139
3140async fn write_all_within<IO>(io: &mut IO, bytes: &[u8], write_timeout: Duration) -> io::Result<()>
3141where
3142 IO: AsyncWrite + Unpin,
3143{
3144 within(write_timeout, io.write_all(bytes))
3145 .await
3146 .map_err(|_| deadline_error("response write"))?
3147}
3148
3149async fn flush_within<IO>(io: &mut IO, write_timeout: Duration) -> io::Result<()>
3150where
3151 IO: AsyncWrite + Unpin,
3152{
3153 within(write_timeout, io.flush())
3154 .await
3155 .map_err(|_| deadline_error("response flush"))?
3156}
3157
3158#[allow(clippy::too_many_arguments)]
3159async fn write_response<IO>(
3160 io: &mut IO,
3161 wire: &mut Vec<u8>,
3162 response: &mut Response,
3163 keep_alive: bool,
3164 send_body: bool,
3165 send_content_length: bool,
3166 write_timeout: Duration,
3167) -> io::Result<()>
3168where
3169 IO: AsyncWrite + Unpin,
3170{
3171 wire.clear();
3172 let streaming = response.is_streaming();
3173 let exact_body_length = response.exact_body_length();
3174 let chunked = streaming && send_body && send_content_length && exact_body_length.is_none();
3175 let content_length = send_content_length.then_some(exact_body_length).flatten();
3176 with_cached_date(|date| {
3177 blazingly_wire::encode_response_head(
3178 wire,
3179 response.status(),
3180 response.headers(),
3181 content_length,
3182 chunked,
3183 keep_alive,
3184 date,
3185 )
3186 })?;
3187 if send_body && !streaming {
3188 wire.extend_from_slice(response.body());
3189 }
3190 write_all_within(io, wire, write_timeout).await?;
3191
3192 if send_body && streaming {
3193 let mut written = 0_u64;
3194 while let Some(chunk) = response.next_body_chunk().await {
3195 let chunk = chunk.map_err(|error| io::Error::other(error.to_string()))?;
3196 if chunk.is_empty() {
3197 continue;
3198 }
3199 let chunk_length = u64::try_from(chunk.len()).unwrap_or(u64::MAX);
3200 written = written
3201 .checked_add(chunk_length)
3202 .ok_or_else(|| io::Error::other("streaming response length overflow"))?;
3203 if exact_body_length.is_some_and(|expected| written > expected) {
3204 return Err(io::Error::other(
3205 "streaming response exceeded its declared exact length",
3206 ));
3207 }
3208 if chunked {
3209 wire.clear();
3210 blazingly_wire::encode_chunk(wire, &chunk)?;
3211 write_all_within(io, wire, write_timeout).await?;
3212 } else {
3213 write_all_within(io, &chunk, write_timeout).await?;
3214 }
3215 }
3216 if exact_body_length.is_some_and(|expected| written != expected) {
3217 return Err(io::Error::other(
3218 "streaming response did not match its declared exact length",
3219 ));
3220 }
3221 if chunked {
3222 write_all_within(io, blazingly_wire::LAST_CHUNK, write_timeout).await?;
3223 }
3224 }
3225 flush_within(io, write_timeout).await
3226}
3227
3228#[allow(clippy::too_many_arguments)]
3229async fn write_response_native(
3230 io: &mut TcpStream,
3231 wire: &mut Vec<u8>,
3232 response: &mut Response,
3233 keep_alive: bool,
3234 send_body: bool,
3235 send_content_length: bool,
3236 write_timeout: Duration,
3237) -> io::Result<()> {
3238 let streaming = response.is_streaming();
3239 let exact_body_length = response.exact_body_length();
3240 let chunked = streaming && send_body && send_content_length && exact_body_length.is_none();
3241 let content_length = send_content_length.then_some(exact_body_length).flatten();
3242 with_cached_date(|date| {
3243 blazingly_wire::encode_response_head(
3244 wire,
3245 response.status(),
3246 response.headers(),
3247 content_length,
3248 chunked,
3249 keep_alive,
3250 date,
3251 )
3252 })?;
3253 if send_body && !streaming {
3254 wire.extend_from_slice(response.body());
3255 }
3256 if streaming {
3257 native_write_all(io, wire, write_timeout).await?;
3258 }
3259
3260 if send_body && streaming {
3261 let mut written = 0_u64;
3262 while let Some(chunk) = response.next_body_chunk().await {
3263 let chunk = chunk.map_err(|error| io::Error::other(error.to_string()))?;
3264 if chunk.is_empty() {
3265 continue;
3266 }
3267 let chunk_length = u64::try_from(chunk.len()).unwrap_or(u64::MAX);
3268 written = written
3269 .checked_add(chunk_length)
3270 .ok_or_else(|| io::Error::other("streaming response length overflow"))?;
3271 if exact_body_length.is_some_and(|expected| written > expected) {
3272 return Err(io::Error::other(
3273 "streaming response exceeded its declared exact length",
3274 ));
3275 }
3276 if chunked {
3277 wire.clear();
3278 blazingly_wire::encode_chunk(wire, &chunk)?;
3279 native_write_all(io, wire, write_timeout).await?;
3280 } else {
3281 let result = within(write_timeout, CompioAsyncWriteExt::write_all(io, chunk))
3282 .await
3283 .map_err(|_| deadline_error("response write"))?;
3284 result.0?;
3285 }
3286 }
3287 if exact_body_length.is_some_and(|expected| written != expected) {
3288 return Err(io::Error::other(
3289 "streaming response did not match its declared exact length",
3290 ));
3291 }
3292 if chunked {
3293 let result = within(
3294 write_timeout,
3295 CompioAsyncWriteExt::write_all(io, blazingly_wire::LAST_CHUNK),
3296 )
3297 .await
3298 .map_err(|_| deadline_error("response write"))?;
3299 result.0?;
3300 }
3301 }
3302 Ok(())
3303}
3304
3305async fn native_write_all(
3306 io: &mut TcpStream,
3307 wire: &mut Vec<u8>,
3308 write_timeout: Duration,
3309) -> io::Result<()> {
3310 if wire.is_empty() {
3311 return Ok(());
3312 }
3313 let result = within(
3314 write_timeout,
3315 CompioAsyncWriteExt::write_all(io, std::mem::take(wire)),
3316 )
3317 .await
3318 .map_err(|_| deadline_error("response write"))?;
3319 let outcome = result.0;
3320 *wire = result.1;
3321 if outcome.is_ok() {
3322 wire.clear();
3323 }
3324 outcome
3325}
3326
3327async fn flush_native_pending(
3328 io: &mut TcpStream,
3329 wire: &mut Vec<u8>,
3330 write_timeout: Duration,
3331) -> io::Result<()> {
3332 native_write_all(io, wire, write_timeout).await
3333}
3334
3335async fn write_rejection<IO>(
3336 io: &mut IO,
3337 wire: &mut Vec<u8>,
3338 status: u16,
3339 code: &str,
3340 message: &str,
3341 write_timeout: Duration,
3342) -> io::Result<()>
3343where
3344 IO: AsyncWrite + Unpin,
3345{
3346 let body = format!(r#"{{"error":{{"code":"{code}","message":"{message}"}}}}"#);
3347 wire.clear();
3348 write!(
3349 wire,
3350 "HTTP/1.1 {status} {}\r\ncontent-type: application/json\r\ncontent-length: {}\r\n",
3351 reason_phrase(status),
3352 body.len()
3353 )?;
3354 wire.extend_from_slice(b"date: ");
3355 with_cached_date(|date| wire.extend_from_slice(date.as_bytes()));
3356 wire.extend_from_slice(b"\r\n");
3357 wire.extend_from_slice(b"connection: close\r\n\r\n");
3358 wire.extend_from_slice(body.as_bytes());
3359 write_all_within(io, wire, write_timeout).await?;
3360 flush_within(io, write_timeout).await
3361}
3362
3363async fn write_rejection_native(
3364 io: &mut TcpStream,
3365 wire: &mut Vec<u8>,
3366 status: u16,
3367 code: &str,
3368 message: &str,
3369 write_timeout: Duration,
3370) -> io::Result<()> {
3371 flush_native_pending(io, wire, write_timeout).await?;
3372 let body = format!(r#"{{"error":{{"code":"{code}","message":"{message}"}}}}"#);
3373 wire.clear();
3374 write!(
3375 wire,
3376 "HTTP/1.1 {status} {}\r\ncontent-type: application/json\r\ncontent-length: {}\r\n",
3377 reason_phrase(status),
3378 body.len()
3379 )?;
3380 wire.extend_from_slice(b"date: ");
3381 with_cached_date(|date| wire.extend_from_slice(date.as_bytes()));
3382 wire.extend_from_slice(b"\r\n");
3383 wire.extend_from_slice(b"connection: close\r\n\r\n");
3384 wire.extend_from_slice(body.as_bytes());
3385 native_write_all(io, wire, write_timeout).await
3386}
3387
3388fn consume_prefix(buffer: &mut Vec<u8>, consumed: usize) {
3389 if consumed == buffer.len() {
3390 buffer.clear();
3391 } else {
3392 buffer.copy_within(consumed.., 0);
3393 buffer.truncate(buffer.len() - consumed);
3394 }
3395}
3396
3397fn with_cached_date<Result>(callback: impl FnOnce(&str) -> Result) -> Result {
3398 let generation = DATE_GENERATION.load(Ordering::Relaxed);
3401 DATE_CACHE.with(|cache| {
3402 let mut cache = cache.borrow_mut();
3403 if cache.generation != generation {
3404 let value = DATE_VALUE
3405 .lock()
3406 .unwrap_or_else(std::sync::PoisonError::into_inner);
3407 cache.value.clone_from(&value);
3408 cache.generation = generation;
3409 }
3410 callback(&cache.value)
3411 })
3412}
3413
3414fn ensure_date_updater() {
3415 DATE_UPDATER.call_once(|| {
3416 refresh_cached_date();
3417 std::thread::Builder::new()
3418 .name("blazingly-date".to_owned())
3419 .spawn(|| {
3420 loop {
3421 std::thread::sleep(Duration::from_secs(1));
3422 refresh_cached_date();
3423 }
3424 })
3425 .expect("failed to start the HTTP Date updater");
3426 });
3427}
3428
3429fn refresh_cached_date() {
3430 let value = httpdate::fmt_http_date(std::time::SystemTime::now());
3431 let mut cached = DATE_VALUE
3432 .lock()
3433 .unwrap_or_else(std::sync::PoisonError::into_inner);
3434 *cached = value;
3435 DATE_GENERATION.fetch_add(1, Ordering::Relaxed);
3436}
3437
3438impl fmt::Debug for Server {
3439 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
3440 formatter
3441 .debug_struct("Server")
3442 .field("limits", &self.limits)
3443 .finish_non_exhaustive()
3444 }
3445}
3446
3447#[cfg(test)]
3448mod tests {
3449 use super::{
3450 HttpMiddleware, MulticoreServer, Rejection, Runtime, Server, ServerLimits, parse_head,
3451 };
3452 use blazingly_core::{
3453 ApiError, HttpMethod, HttpUpgrade, InputDescriptor, InputSource, Json, MultipartError,
3454 OperationDescriptor, ResponseDescriptor, ResponseHeader, SchemaKind, TypeDescriptor,
3455 };
3456 use blazingly_executor::{
3457 DependencyError, ExecutableApp, ExecutableOperation, ExecutionOutcome, FromInvocation,
3458 OperationFuture, Plugin, UploadBody,
3459 };
3460 use blazingly_http::{HttpRequestContext, Request, Response};
3461 use compio::io::{
3462 AsyncReadExt as CompioAsyncReadExt, AsyncWrite as CompioAsyncWrite,
3463 AsyncWriteExt as CompioAsyncWriteExt,
3464 };
3465 use futures_lite::future;
3466 use futures_lite::io::{AsyncRead, AsyncWrite};
3467 use std::cell::RefCell;
3468 use std::io;
3469 use std::num::NonZeroUsize;
3470 use std::pin::Pin;
3471 use std::rc::Rc;
3472 use std::sync::Arc;
3473 use std::sync::atomic::Ordering;
3474 use std::task::{Context, Poll};
3475 use std::time::Duration;
3476
3477 struct ScriptedTransport {
3480 input: Vec<u8>,
3481 close_after_input: bool,
3482 written: Rc<RefCell<Vec<u8>>>,
3483 }
3484
3485 impl ScriptedTransport {
3486 fn closing(input: &[u8]) -> Self {
3487 Self {
3488 input: input.to_vec(),
3489 close_after_input: true,
3490 written: Rc::new(RefCell::new(Vec::new())),
3491 }
3492 }
3493
3494 fn half_open(input: &[u8]) -> Self {
3495 Self {
3496 input: input.to_vec(),
3497 close_after_input: false,
3498 written: Rc::new(RefCell::new(Vec::new())),
3499 }
3500 }
3501
3502 fn output(&self) -> Rc<RefCell<Vec<u8>>> {
3504 Rc::clone(&self.written)
3505 }
3506
3507 fn written(&self) -> String {
3508 String::from_utf8_lossy(&self.written.borrow()).into_owned()
3509 }
3510 }
3511
3512 impl AsyncRead for ScriptedTransport {
3513 fn poll_read(
3514 mut self: Pin<&mut Self>,
3515 _context: &mut Context<'_>,
3516 buffer: &mut [u8],
3517 ) -> Poll<io::Result<usize>> {
3518 if self.input.is_empty() {
3519 return if self.close_after_input {
3520 Poll::Ready(Ok(0))
3521 } else {
3522 Poll::Pending
3523 };
3524 }
3525 let read = self.input.len().min(buffer.len());
3526 buffer[..read].copy_from_slice(&self.input[..read]);
3527 self.input.drain(..read);
3528 Poll::Ready(Ok(read))
3529 }
3530 }
3531
3532 impl AsyncWrite for ScriptedTransport {
3533 fn poll_write(
3534 self: Pin<&mut Self>,
3535 _context: &mut Context<'_>,
3536 buffer: &[u8],
3537 ) -> Poll<io::Result<usize>> {
3538 self.written.borrow_mut().extend_from_slice(buffer);
3539 Poll::Ready(Ok(buffer.len()))
3540 }
3541
3542 fn poll_flush(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
3543 Poll::Ready(Ok(()))
3544 }
3545
3546 fn poll_close(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
3547 Poll::Ready(Ok(()))
3548 }
3549 }
3550
3551 thread_local! {
3552 static EVENTS: RefCell<Vec<String>> = const { RefCell::new(Vec::new()) };
3553 }
3554
3555 fn record(event: impl Into<String>) {
3557 EVENTS.with(|events| events.borrow_mut().push(event.into()));
3558 }
3559
3560 fn recorded_events() -> Vec<String> {
3561 EVENTS.with(|events| events.borrow().clone())
3562 }
3563
3564 fn clear_events() {
3565 EVENTS.with(|events| events.borrow_mut().clear());
3566 }
3567
3568 fn event_position(events: &[String], event: &str) -> usize {
3569 events
3570 .iter()
3571 .position(|recorded| recorded == event)
3572 .unwrap_or_else(|| panic!("{event} was never recorded: {events:?}"))
3573 }
3574
3575 struct StagedTransport {
3581 segments: Vec<(&'static str, Vec<u8>)>,
3582 next: usize,
3583 written: Rc<RefCell<Vec<u8>>>,
3584 }
3585
3586 impl StagedTransport {
3587 fn new(segments: Vec<(&'static str, Vec<u8>)>) -> Self {
3588 Self {
3589 segments,
3590 next: 0,
3591 written: Rc::new(RefCell::new(Vec::new())),
3592 }
3593 }
3594
3595 fn written(&self) -> String {
3596 String::from_utf8_lossy(&self.written.borrow()).into_owned()
3597 }
3598 }
3599
3600 impl AsyncRead for StagedTransport {
3601 fn poll_read(
3602 mut self: Pin<&mut Self>,
3603 _context: &mut Context<'_>,
3604 buffer: &mut [u8],
3605 ) -> Poll<io::Result<usize>> {
3606 let Some((label, bytes)) = self.segments.get(self.next).cloned() else {
3607 return Poll::Ready(Ok(0));
3608 };
3609 assert!(
3610 bytes.len() <= buffer.len(),
3611 "a staged segment must fit one read"
3612 );
3613 self.next += 1;
3614 record(format!("client:{label}"));
3615 buffer[..bytes.len()].copy_from_slice(&bytes);
3616 Poll::Ready(Ok(bytes.len()))
3617 }
3618 }
3619
3620 impl AsyncWrite for StagedTransport {
3621 fn poll_write(
3622 self: Pin<&mut Self>,
3623 _context: &mut Context<'_>,
3624 buffer: &[u8],
3625 ) -> Poll<io::Result<usize>> {
3626 if buffer.starts_with(b"HTTP/1.1 100 ") {
3627 record("server:continue");
3628 }
3629 self.written.borrow_mut().extend_from_slice(buffer);
3630 Poll::Ready(Ok(buffer.len()))
3631 }
3632
3633 fn poll_flush(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
3634 Poll::Ready(Ok(()))
3635 }
3636
3637 fn poll_close(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
3638 Poll::Ready(Ok(()))
3639 }
3640 }
3641
3642 struct StampMiddleware;
3643
3644 impl HttpMiddleware for StampMiddleware {
3645 fn on_response(
3646 &self,
3647 _context: &HttpRequestContext<'_>,
3648 _operation: Option<&OperationDescriptor>,
3649 response: &mut Response,
3650 ) {
3651 response.set_header("x-native-middleware", "applied");
3652 }
3653 }
3654
3655 fn ping_operation() -> ExecutableOperation {
3656 let descriptor = OperationDescriptor::new(
3657 HttpMethod::Get,
3658 "/ping",
3659 "ping",
3660 "Ping",
3661 None,
3662 vec![ResponseDescriptor::success(200, None)],
3663 )
3664 .expect("the ping descriptor is valid");
3665 ExecutableOperation::empty(descriptor, || async { Json("pong") })
3666 }
3667
3668 fn ping_app() -> ExecutableApp {
3669 ExecutableApp::new([ping_operation()]).expect("the ping app compiles")
3670 }
3671
3672 fn echo_upgrade_operation() -> ExecutableOperation {
3675 let descriptor = OperationDescriptor::new(
3676 HttpMethod::Get,
3677 "/upgrade",
3678 "upgrade.echo",
3679 "Echo upgrade",
3680 None,
3681 vec![ResponseDescriptor::success(200, None)],
3682 )
3683 .expect("the upgrade descriptor is valid");
3684 ExecutableOperation::empty(descriptor, || async {
3685 HttpUpgrade::new(
3686 "echo",
3687 vec![
3688 ResponseHeader::new("connection", "Upgrade"),
3689 ResponseHeader::new("upgrade", "echo"),
3690 ],
3691 |mut io| {
3692 Box::pin(async move {
3693 while let Some(bytes) = io.read().await? {
3694 if bytes.is_empty() {
3695 continue;
3696 }
3697 io.write(bytes).await?;
3698 break;
3699 }
3700 io.shutdown().await
3701 })
3702 },
3703 )
3704 })
3705 }
3706
3707 fn upload_operation(
3712 path: &'static str,
3713 id: &'static str,
3714 read_body: bool,
3715 ) -> ExecutableOperation {
3716 let descriptor = OperationDescriptor::new(
3717 HttpMethod::Post,
3718 path,
3719 id,
3720 "Upload",
3721 None,
3722 vec![ResponseDescriptor::success(200, None)],
3723 )
3724 .expect("the upload descriptor is valid")
3725 .with_inputs(vec![InputDescriptor::new(
3726 "body",
3727 InputSource::Stream,
3728 true,
3729 TypeDescriptor::scalar("UploadBody", SchemaKind::Binary),
3730 )]);
3731 ExecutableOperation::typed(descriptor, move |input| {
3732 let mut body = UploadBody::from_invocation(&input, "body", true)?;
3733 Ok(Box::pin(async move {
3734 if !read_body {
3735 std::future::pending::<()>().await;
3736 }
3737 let mut bytes = 0_usize;
3738 while let Some(chunk) = body.next_chunk().await {
3739 match chunk {
3740 Ok(chunk) => bytes += chunk.len(),
3741 Err(error) => {
3742 return ExecutionOutcome::InternalError {
3743 code: error.code,
3744 message: error.message,
3745 };
3746 }
3747 }
3748 }
3749 ExecutionOutcome::Success {
3750 status: 200,
3751 headers: vec![ResponseHeader::new("content-type", "text/plain")],
3752 body: Some(format!("bytes={bytes}").into_bytes()),
3753 background: Vec::new(),
3754 }
3755 }) as OperationFuture)
3756 })
3757 }
3758
3759 fn logging_upload_operation() -> ExecutableOperation {
3761 let descriptor = OperationDescriptor::new(
3762 HttpMethod::Post,
3763 "/staged-upload",
3764 "upload.staged",
3765 "Staged upload",
3766 None,
3767 vec![ResponseDescriptor::success(200, None)],
3768 )
3769 .expect("the staged upload descriptor is valid")
3770 .with_inputs(vec![InputDescriptor::new(
3771 "body",
3772 InputSource::Stream,
3773 true,
3774 TypeDescriptor::scalar("UploadBody", SchemaKind::Binary),
3775 )]);
3776 ExecutableOperation::typed(descriptor, |input| {
3777 let mut body = UploadBody::from_invocation(&input, "body", true)?;
3778 Ok(Box::pin(async move {
3779 let mut bytes = 0_usize;
3780 while let Some(chunk) = body.next_chunk().await {
3781 match chunk {
3782 Ok(chunk) => {
3783 record(format!("handler:{}", String::from_utf8_lossy(&chunk)));
3784 bytes += chunk.len();
3785 }
3786 Err(error) => {
3787 return ExecutionOutcome::InternalError {
3788 code: error.code,
3789 message: error.message,
3790 };
3791 }
3792 }
3793 }
3794 ExecutionOutcome::Success {
3795 status: 200,
3796 headers: vec![ResponseHeader::new("content-type", "text/plain")],
3797 body: Some(format!("bytes={bytes}").into_bytes()),
3798 background: Vec::new(),
3799 }
3800 }) as OperationFuture)
3801 })
3802 }
3803
3804 const MULTIPART_BOUNDARY: &str = "blazingly-native-cover";
3805
3806 fn multipart_document(data: &[u8]) -> Vec<u8> {
3808 let mut body = format!(
3809 "--{MULTIPART_BOUNDARY}\r\nContent-Disposition: form-data; name=\"file\"; \
3810 filename=\"cover.jpg\"\r\nContent-Type: image/jpeg\r\n\r\n"
3811 )
3812 .into_bytes();
3813 body.extend_from_slice(data);
3814 body.extend_from_slice(format!("\r\n--{MULTIPART_BOUNDARY}--\r\n").as_bytes());
3815 body
3816 }
3817
3818 fn multipart_request(body: &[u8]) -> Vec<u8> {
3820 let mut request = format!(
3821 "POST /cover HTTP/1.1\r\nhost: localhost\r\ncontent-type: multipart/form-data; \
3822 boundary={MULTIPART_BOUNDARY}\r\ncontent-length: {}\r\nconnection: close\r\n\r\n",
3823 body.len()
3824 )
3825 .into_bytes();
3826 request.extend_from_slice(body);
3827 request
3828 }
3829
3830 fn multipart_failure(error: &MultipartError) -> ExecutionOutcome {
3831 error.clone().into_failure().map_or_else(
3832 |error| ExecutionOutcome::InternalError {
3833 code: error.code,
3834 message: error.message,
3835 },
3836 ExecutionOutcome::DomainError,
3837 )
3838 }
3839
3840 fn streaming_multipart_operation() -> ExecutableOperation {
3845 let descriptor = OperationDescriptor::new(
3846 HttpMethod::Post,
3847 "/cover",
3848 "cover.stream",
3849 "Streaming cover upload",
3850 None,
3851 vec![ResponseDescriptor::success(200, None)],
3852 )
3853 .expect("the streaming multipart descriptor is valid")
3854 .with_inputs(vec![InputDescriptor::new(
3855 "body",
3856 InputSource::Stream,
3857 true,
3858 TypeDescriptor::scalar("UploadBody", SchemaKind::Binary),
3859 )]);
3860 ExecutableOperation::typed(descriptor, |input| {
3861 let body = UploadBody::from_invocation(&input, "body", true)?;
3862 Ok(Box::pin(async move {
3863 let mut multipart = match body.into_multipart() {
3864 Ok(multipart) => multipart,
3865 Err(error) => return multipart_failure(&error),
3866 };
3867 let mut bytes = 0_usize;
3868 loop {
3869 let mut field = match multipart.next_field().await {
3870 Ok(Some(field)) => field,
3871 Ok(None) => break,
3872 Err(error) => return multipart_failure(&error),
3873 };
3874 if field.name() != "file" {
3875 continue;
3876 }
3877 loop {
3878 match field.next_chunk().await {
3879 Ok(Some(chunk)) => {
3880 record(format!("handler:{}", chunk.len()));
3881 bytes += chunk.len();
3882 }
3883 Ok(None) => break,
3884 Err(error) => return multipart_failure(&error),
3885 }
3886 }
3887 }
3888 ExecutionOutcome::Success {
3889 status: 200,
3890 headers: vec![ResponseHeader::new("content-type", "text/plain")],
3891 body: Some(format!("bytes={bytes}").into_bytes()),
3892 background: Vec::new(),
3893 }
3894 }) as OperationFuture)
3895 })
3896 }
3897
3898 #[test]
3899 fn the_plaintext_socket_streams_a_multipart_upload() {
3900 let payload = vec![7_u8; 1 << 20];
3901 let response = native_exchange_with_limits(
3902 vec![streaming_multipart_operation()],
3903 multipart_request(&multipart_document(&payload)),
3904 ServerLimits::new().with_max_body_bytes(4 * 1024 * 1024),
3905 );
3906
3907 let response = String::from_utf8_lossy(&response).into_owned();
3910 let head = response.lines().next().unwrap_or_default().to_owned();
3911 assert!(response.starts_with("HTTP/1.1 200 "), "{head}");
3912 assert!(
3913 response.ends_with(&format!("bytes={}", payload.len())),
3914 "{head}"
3915 );
3916 }
3917
3918 #[test]
3919 fn the_compatibility_transport_streams_a_multipart_upload_before_the_client_finishes_it() {
3920 clear_events();
3921 let document = multipart_document(&[3_u8; 32]);
3922 let head = format!(
3923 "POST /cover HTTP/1.1\r\nhost: localhost\r\ncontent-type: multipart/form-data; \
3924 boundary={MULTIPART_BOUNDARY}\r\ncontent-length: {}\r\nconnection: close\r\n\r\n",
3925 document.len()
3926 )
3927 .into_bytes();
3928 let split = document.len() - 16;
3931 let server = Server::new(
3932 ExecutableApp::new([streaming_multipart_operation()])
3933 .expect("the streaming multipart app compiles"),
3934 );
3935 let mut transport = StagedTransport::new(vec![
3936 ("head", head),
3937 ("first", document[..split].to_vec()),
3938 ("last", document[split..].to_vec()),
3939 ]);
3940
3941 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
3942
3943 let events = recorded_events();
3944 let first_chunk = events
3945 .iter()
3946 .position(|event| event.starts_with("handler:"))
3947 .unwrap_or_else(|| panic!("the handler never saw a chunk: {events:?}"));
3948 assert!(
3949 first_chunk < event_position(&events, "client:last"),
3950 "the compatibility transport buffered the whole body before dispatch: {events:?}"
3951 );
3952 let written = transport.written();
3953 assert!(written.starts_with("HTTP/1.1 200 "), "{written}");
3954 assert!(written.ends_with("bytes=32"), "{written}");
3955 }
3956
3957 #[test]
3958 fn a_malformed_multipart_body_is_rejected_while_streaming() {
3959 let response = native_exchange(
3960 vec![streaming_multipart_operation()],
3961 multipart_request(b"this is not a multipart document"),
3962 );
3963
3964 let response = String::from_utf8_lossy(&response).into_owned();
3965 assert!(response.starts_with("HTTP/1.1 422 "), "{response}");
3966 assert!(response.contains("invalid_multipart"), "{response}");
3967 }
3968
3969 #[test]
3970 fn a_truncated_multipart_body_is_not_answered_with_a_success() {
3971 let document = multipart_document(&[5_u8; 4096]);
3972 let mut request = multipart_request(&document);
3973 request.truncate(request.len() - 2048);
3976 let response = native_exchange(vec![streaming_multipart_operation()], request);
3977
3978 let response = String::from_utf8_lossy(&response).into_owned();
3979 assert!(!response.contains(" 200 "), "{response}");
3980 assert!(response.starts_with("HTTP/1.1 400 "), "{response}");
3981 assert!(response.contains("incomplete_body"), "{response}");
3982 }
3983
3984 #[test]
3985 fn transport_limits_still_bound_a_streamed_multipart_upload() {
3986 let document = multipart_document(&[1_u8; 64]);
3987 let mut request = format!(
3988 "POST /cover HTTP/1.1\r\nhost: localhost\r\ncontent-type: multipart/form-data; \
3989 boundary={MULTIPART_BOUNDARY}\r\ntransfer-encoding: chunked\r\nconnection: \
3990 close\r\n\r\n"
3991 )
3992 .into_bytes();
3993 for piece in document.chunks(32) {
3994 request.extend_from_slice(format!("{:x}\r\n", piece.len()).as_bytes());
3995 request.extend_from_slice(piece);
3996 request.extend_from_slice(b"\r\n");
3997 }
3998 request.extend_from_slice(b"0\r\n\r\n");
3999
4000 let response = native_exchange_with_limits(
4001 vec![streaming_multipart_operation()],
4002 request,
4003 ServerLimits::new().with_max_chunks(2),
4004 );
4005
4006 let response = String::from_utf8_lossy(&response).into_owned();
4007 assert!(response.starts_with("HTTP/1.1 413 "), "{response}");
4008 assert!(response.contains("too_many_chunks"), "{response}");
4009 }
4010
4011 fn buffered_post_operation() -> ExecutableOperation {
4014 let descriptor = OperationDescriptor::new(
4015 HttpMethod::Post,
4016 "/buffered",
4017 "buffered.post",
4018 "Buffered post",
4019 None,
4020 vec![ResponseDescriptor::success(200, None)],
4021 )
4022 .expect("the buffered descriptor is valid");
4023 ExecutableOperation::empty(descriptor, || async { Json("stored") })
4024 }
4025
4026 #[cfg(feature = "http2")]
4027 fn never_completing_operation() -> ExecutableOperation {
4028 let descriptor = OperationDescriptor::new(
4029 HttpMethod::Get,
4030 "/forever",
4031 "forever.wait",
4032 "Never completes",
4033 None,
4034 vec![ResponseDescriptor::success(200, None)],
4035 )
4036 .expect("the waiting descriptor is valid");
4037 ExecutableOperation::empty(descriptor, || async {
4038 std::future::pending::<()>().await;
4039 blazingly_core::NoContent
4040 })
4041 }
4042
4043 fn native_exchange(operations: Vec<ExecutableOperation>, request: Vec<u8>) -> Vec<u8> {
4046 native_exchange_with_limits(operations, request, ServerLimits::new())
4047 }
4048
4049 fn native_exchange_with_limits(
4051 operations: Vec<ExecutableOperation>,
4052 request: Vec<u8>,
4053 limits: ServerLimits,
4054 ) -> Vec<u8> {
4055 let runtime = Runtime::new().expect("the Compio runtime starts");
4056 runtime.block_on(async move {
4057 let listener = compio::net::TcpListener::bind("127.0.0.1:0")
4058 .await
4059 .expect("the loopback listener binds");
4060 let address = listener.local_addr().expect("the listener has an address");
4061 let app = Rc::new(super::HttpApp::new(
4062 ExecutableApp::new(operations).expect("the test app compiles"),
4063 ));
4064 let mut client = compio::net::TcpStream::connect(address)
4065 .await
4066 .expect("the client connects");
4067 let (stream, peer) = listener.accept().await.expect("the server accepts");
4068 let served = super::spawn(async move {
4069 super::serve_accepted(
4070 app.as_ref(),
4071 stream,
4072 peer,
4073 None,
4074 super::ConnectionSetup {
4075 limits,
4076 request_timeout: None,
4077 #[cfg(feature = "tls")]
4078 tls_acceptor: None,
4079 },
4080 )
4081 .await;
4082 });
4083 let write = CompioAsyncWriteExt::write_all(&mut client, request).await;
4084 write.0.expect("the request is written");
4085 CompioAsyncWrite::shutdown(&mut client)
4086 .await
4087 .expect("the client half-closes");
4088 let read = CompioAsyncReadExt::read_to_end(&mut client, Vec::new()).await;
4089 read.0.expect("the response is read");
4090 served.await.expect("the connection task finishes");
4091 read.1
4092 })
4093 }
4094
4095 #[test]
4096 fn header_parser_honors_limits_below_the_inline_capacity() {
4097 let request = b"GET / HTTP/1.1\r\nhost: localhost\r\nx-extra: value\r\n\r\n";
4098 let result = parse_head(request, ServerLimits::new().with_max_headers(1));
4099
4100 assert!(matches!(result, Err(Rejection { status: 431, .. })));
4101 }
4102
4103 #[test]
4104 fn registered_middleware_runs_on_a_native_connection() {
4105 let server = Server::new(ping_app()).with_middleware(StampMiddleware);
4106 let mut transport =
4107 ScriptedTransport::closing(b"GET /ping HTTP/1.1\r\nhost: localhost\r\n\r\n");
4108
4109 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4110
4111 let written = transport.written();
4112 assert!(written.starts_with("HTTP/1.1 200 "), "{written}");
4113 assert!(
4114 written.contains("x-native-middleware: applied"),
4115 "{written}"
4116 );
4117 }
4118
4119 #[test]
4120 fn shared_middleware_runs_on_a_native_connection() {
4121 let server = Server::new(ping_app()).with_shared_middleware(Rc::new(StampMiddleware));
4122 let mut transport =
4123 ScriptedTransport::closing(b"GET /ping HTTP/1.1\r\nhost: localhost\r\n\r\n");
4124
4125 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4126
4127 assert!(
4128 transport.written().contains("x-native-middleware: applied"),
4129 "{}",
4130 transport.written()
4131 );
4132 }
4133
4134 #[test]
4135 fn header_read_deadline_answers_408_and_closes_a_half_open_request() {
4136 let limits = ServerLimits::new().with_header_read_timeout(Duration::from_millis(50));
4137 let server = Server::new(ping_app()).with_limits(limits);
4138 let runtime = Runtime::new().expect("the Compio runtime starts");
4139
4140 let written = runtime.block_on(async {
4141 let mut transport = ScriptedTransport::half_open(b"GET /ping HTTP/1.1\r\nhost: loc");
4142 server
4143 .serve_io(&mut transport)
4144 .await
4145 .expect("an expired header deadline closes the connection cleanly");
4146 transport.written()
4147 });
4148
4149 assert!(written.starts_with("HTTP/1.1 408 "), "{written}");
4150 assert!(written.contains("request_timeout"), "{written}");
4151 }
4152
4153 struct StalledWriter {
4156 input: Vec<u8>,
4157 }
4158
4159 impl AsyncRead for StalledWriter {
4160 fn poll_read(
4161 mut self: Pin<&mut Self>,
4162 _context: &mut Context<'_>,
4163 buffer: &mut [u8],
4164 ) -> Poll<io::Result<usize>> {
4165 if self.input.is_empty() {
4166 return Poll::Pending;
4167 }
4168 let read = self.input.len().min(buffer.len());
4169 buffer[..read].copy_from_slice(&self.input[..read]);
4170 self.input.drain(..read);
4171 Poll::Ready(Ok(read))
4172 }
4173 }
4174
4175 impl AsyncWrite for StalledWriter {
4176 fn poll_write(
4177 self: Pin<&mut Self>,
4178 _context: &mut Context<'_>,
4179 _buffer: &[u8],
4180 ) -> Poll<io::Result<usize>> {
4181 Poll::Pending
4182 }
4183
4184 fn poll_flush(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
4185 Poll::Pending
4186 }
4187
4188 fn poll_close(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<io::Result<()>> {
4189 Poll::Ready(Ok(()))
4190 }
4191 }
4192
4193 #[test]
4194 fn write_deadline_abandons_a_peer_that_never_drains_the_response() {
4195 let limits = ServerLimits::new().with_write_timeout(Duration::from_millis(50));
4199 let server = Server::new(ping_app()).with_limits(limits);
4200 let runtime = Runtime::new().expect("the Compio runtime starts");
4201
4202 let error = runtime.block_on(async {
4203 let mut transport = StalledWriter {
4204 input: b"GET /ping HTTP/1.1\r\nhost: localhost\r\n\r\n".to_vec(),
4205 };
4206 server
4207 .serve_io(&mut transport)
4208 .await
4209 .expect_err("a stalled write must expire instead of hanging")
4210 });
4211
4212 assert_eq!(error.kind(), io::ErrorKind::TimedOut, "{error}");
4213 }
4214
4215 #[test]
4216 fn idle_deadline_closes_a_silent_connection_without_a_response() {
4217 let limits = ServerLimits::new().with_idle_timeout(Duration::from_millis(50));
4218 let server = Server::new(ping_app()).with_limits(limits);
4219 let runtime = Runtime::new().expect("the Compio runtime starts");
4220
4221 let written = runtime.block_on(async {
4222 let mut transport = ScriptedTransport::half_open(b"");
4223 server
4224 .serve_io(&mut transport)
4225 .await
4226 .expect("an expired idle deadline closes the connection cleanly");
4227 transport.written()
4228 });
4229
4230 assert!(written.is_empty(), "{written}");
4231 }
4232
4233 #[test]
4234 fn multicore_serve_refuses_to_boot_when_worker_startup_fails() {
4235 let workers = NonZeroUsize::new(1).expect("one worker is non-zero");
4236 let server = MulticoreServer::new(workers, || {
4237 ExecutableApp::from_plugin(Plugin::new("app").routes([ping_operation()]).on_startup(
4238 || async {
4239 Err(DependencyError::internal(
4240 "startup_failed",
4241 "startup hook refused to boot",
4242 ))
4243 },
4244 ))
4245 .expect("the failing app compiles")
4246 });
4247
4248 let error = server
4249 .serve("127.0.0.1:0")
4250 .expect_err("a failing startup hook aborts serve");
4251
4252 assert_eq!(error.to_string(), "startup hook refused to boot");
4253 }
4254
4255 #[test]
4256 fn owned_compatibility_transport_hands_ownership_to_an_upgrade() {
4257 let server = Server::new(
4258 ExecutableApp::new([echo_upgrade_operation()]).expect("the upgrade app compiles"),
4259 );
4260 let transport = ScriptedTransport::closing(
4261 b"GET /upgrade HTTP/1.1\r\nhost: localhost\r\nconnection: Upgrade\r\nupgrade: echo\r\n\r\nPING",
4262 );
4263 let written = transport.output();
4264
4265 future::block_on(server.serve_owned_io(transport)).expect("the upgraded session completes");
4266
4267 let written = String::from_utf8_lossy(&written.borrow()).into_owned();
4268 assert!(
4269 written.starts_with("HTTP/1.1 101 Switching Protocols\r\n"),
4270 "{written}"
4271 );
4272 assert!(written.contains("upgrade: echo\r\n"), "{written}");
4273 assert!(
4274 written.ends_with("PING"),
4275 "the bytes buffered behind the handshake were lost: {written}"
4276 );
4277 }
4278
4279 #[test]
4280 fn borrowed_transport_still_refuses_an_upgrade() {
4281 let server = Server::new(
4282 ExecutableApp::new([echo_upgrade_operation()]).expect("the upgrade app compiles"),
4283 );
4284 let mut transport = ScriptedTransport::closing(
4285 b"GET /upgrade HTTP/1.1\r\nhost: localhost\r\nconnection: Upgrade\r\nupgrade: echo\r\n\r\n",
4286 );
4287
4288 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4289
4290 let written = transport.written();
4291 assert!(written.starts_with("HTTP/1.1 501 "), "{written}");
4292 assert!(
4293 written.contains("upgrade_transport_unsupported"),
4294 "{written}"
4295 );
4296 }
4297
4298 #[cfg(any(feature = "http2", feature = "tls"))]
4299 #[test]
4300 fn compatibility_stream_upgrade_matches_the_tls_transport() {
4301 let runtime = Runtime::new().expect("the Compio runtime starts");
4302 let response = runtime.block_on(async {
4303 let listener = compio::net::TcpListener::bind("127.0.0.1:0")
4304 .await
4305 .expect("the loopback listener binds");
4306 let address = listener.local_addr().expect("the listener has an address");
4307 let app = Rc::new(super::HttpApp::new(
4308 ExecutableApp::new([echo_upgrade_operation()]).expect("the upgrade app compiles"),
4309 ));
4310 let mut client = compio::net::TcpStream::connect(address)
4311 .await
4312 .expect("the client connects");
4313 let (stream, peer) = listener.accept().await.expect("the server accepts");
4314 let served = super::spawn(async move {
4315 super::serve_compat_connection(
4316 app.as_ref(),
4317 ServerLimits::new(),
4318 Box::pin(super::AsyncStream::new(stream)),
4319 Some(peer),
4320 "http",
4321 None,
4322 None,
4323 )
4324 .await
4325 });
4326 let request = b"GET /upgrade HTTP/1.1\r\nhost: localhost\r\nconnection: Upgrade\r\nupgrade: echo\r\n\r\nPING".to_vec();
4327 let write = CompioAsyncWriteExt::write_all(&mut client, request).await;
4328 write.0.expect("the request is written");
4329 let read = CompioAsyncReadExt::read_to_end(&mut client, Vec::new()).await;
4330 read.0.expect("the response is read");
4331 served
4332 .await
4333 .expect("the connection task finishes")
4334 .expect("the upgraded session completes");
4335 read.1
4336 });
4337
4338 let response = String::from_utf8_lossy(&response).into_owned();
4339 assert!(
4340 response.starts_with("HTTP/1.1 101 Switching Protocols\r\n"),
4341 "{response}"
4342 );
4343 assert!(
4344 response.ends_with("PING"),
4345 "the compatibility transport lost the buffered frame: {response}"
4346 );
4347 }
4348
4349 #[test]
4350 fn native_socket_upgrade_survives_every_feature_combination() {
4351 let response = native_exchange(
4352 vec![echo_upgrade_operation()],
4353 b"GET /upgrade HTTP/1.1\r\nhost: localhost\r\nconnection: Upgrade\r\nupgrade: echo\r\n\r\nPING".to_vec(),
4354 );
4355
4356 let response = String::from_utf8_lossy(&response).into_owned();
4357 assert!(
4358 response.starts_with("HTTP/1.1 101 Switching Protocols\r\n"),
4359 "{response}"
4360 );
4361 assert!(response.ends_with("PING"), "{response}");
4362 }
4363
4364 #[test]
4365 fn native_socket_streams_uploads_in_every_feature_combination() {
4366 let response = native_exchange(
4367 vec![upload_operation("/upload", "upload.consume", true)],
4368 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 11\r\nconnection: close\r\n\r\nhello world"
4369 .to_vec(),
4370 );
4371
4372 let response = String::from_utf8_lossy(&response).into_owned();
4373 assert!(response.starts_with("HTTP/1.1 200 OK\r\n"), "{response}");
4374 assert!(
4375 response.ends_with("bytes=11"),
4376 "the streaming upload seam did not run: {response}"
4377 );
4378 }
4379
4380 #[cfg(feature = "http2")]
4381 fn h2_client(
4382 method: &str,
4383 path: &str,
4384 body: Option<&[u8]>,
4385 ) -> (
4386 shiguredo_http2::Connection,
4387 shiguredo_http2::StreamId,
4388 Vec<u8>,
4389 ) {
4390 use shiguredo_http2::{Connection, HeaderField, Limits};
4391
4392 let mut client = Connection::client(Limits::default());
4393 client.initiate().expect("the client preface is written");
4394 let mut headers = vec![
4395 HeaderField::new(":method", method).expect("method header"),
4396 HeaderField::new(":path", path).expect("path header"),
4397 HeaderField::new(":scheme", "http").expect("scheme header"),
4398 HeaderField::new(":authority", "localhost").expect("authority header"),
4399 ];
4400 if let Some(body) = body {
4401 headers.push(
4402 HeaderField::new("content-length", body.len().to_string())
4403 .expect("content-length header"),
4404 );
4405 }
4406 let stream = client
4407 .start_stream(headers, body.is_none())
4408 .expect("the request stream starts");
4409 if let Some(body) = body {
4410 client
4411 .send_data(stream, body.to_vec(), true)
4412 .expect("the request body is sent");
4413 }
4414 let mut bytes = Vec::new();
4415 while let Some(output) = client.poll_output() {
4416 bytes.extend_from_slice(&output);
4417 }
4418 (client, stream, bytes)
4419 }
4420
4421 #[cfg(feature = "http2")]
4422 fn h2_events(
4423 client: &mut shiguredo_http2::Connection,
4424 output: &[u8],
4425 ) -> Vec<shiguredo_http2::Event> {
4426 client.feed(output).expect("the server output is accepted");
4427 client.process().expect("the server output is processed");
4428 let mut events = Vec::new();
4429 while let Some(event) = client.poll_event() {
4430 events.push(event);
4431 }
4432 events
4433 }
4434
4435 #[cfg(feature = "http2")]
4436 #[test]
4437 fn http2_preface_still_reaches_the_codec_on_a_plaintext_socket() {
4438 use shiguredo_http2::Event;
4439
4440 let (mut client, stream, request) = h2_client("GET", "/ping", None);
4441 let response = native_exchange(vec![ping_operation()], request);
4442 let events = h2_events(&mut client, &response);
4443
4444 let status = events.iter().find_map(|event| match event {
4445 Event::HeadersReceived {
4446 stream_id, headers, ..
4447 } if *stream_id == stream => headers
4448 .iter()
4449 .find(|header| header.name() == b":status")
4450 .map(|header| header.value().to_vec()),
4451 _ => None,
4452 });
4453 assert_eq!(status.as_deref(), Some(b"200".as_slice()));
4454 }
4455
4456 #[cfg(feature = "http2")]
4457 #[test]
4458 fn http2_streams_request_bodies_through_the_upload_seam() {
4459 use shiguredo_http2::Event;
4460
4461 let (mut client, stream, request) = h2_client("POST", "/upload", Some(b"hello world"));
4462 let server = Server::new(
4463 ExecutableApp::new([upload_operation("/upload", "upload.consume", true)])
4464 .expect("the upload app compiles"),
4465 );
4466 let mut transport = ScriptedTransport::closing(&request);
4467 future::block_on(server.serve_http2_io(&mut transport))
4468 .expect("the HTTP/2 exchange completes");
4469 let output = transport.output().borrow().clone();
4470 let events = h2_events(&mut client, &output);
4471
4472 let mut body = Vec::new();
4473 for event in &events {
4474 if let Event::DataReceived {
4475 stream_id, data, ..
4476 } = event
4477 && *stream_id == stream
4478 {
4479 body.extend_from_slice(data);
4480 }
4481 }
4482 assert_eq!(String::from_utf8_lossy(&body), "bytes=11");
4483 }
4484
4485 #[cfg(feature = "http2")]
4486 #[test]
4487 fn http2_returns_receive_window_credit_only_as_the_handler_consumes() {
4488 use shiguredo_http2::StreamId;
4489
4490 let consuming = h2_window_credit(true);
4491 assert!(
4492 consuming.iter().any(|(stream_id, increment)| matches!(
4493 stream_id,
4494 StreamId::Connection
4495 ) && *increment == 11),
4496 "a consuming handler must return exactly the credit it read: {consuming:?}"
4497 );
4498
4499 let idle = h2_window_credit(false);
4500 assert!(
4501 idle.is_empty(),
4502 "credit was returned for bytes the handler never read: {idle:?}"
4503 );
4504 }
4505
4506 #[cfg(feature = "http2")]
4509 fn h2_window_credit(read_body: bool) -> Vec<(shiguredo_http2::StreamId, u32)> {
4510 use shiguredo_http2::Event;
4511
4512 let (mut client, _, request) = h2_client("POST", "/upload", Some(b"hello world"));
4513 let server = Server::new(
4514 ExecutableApp::new([upload_operation("/upload", "upload.consume", read_body)])
4515 .expect("the upload app compiles"),
4516 );
4517 let runtime = Runtime::new().expect("the Compio runtime starts");
4518 let output = runtime.block_on(async {
4519 let mut transport = if read_body {
4520 ScriptedTransport::closing(&request)
4521 } else {
4522 ScriptedTransport::half_open(&request)
4523 };
4524 let recorded = transport.output();
4525 let _ = compio::time::timeout(
4526 Duration::from_millis(500),
4527 server.serve_http2_io(&mut transport),
4528 )
4529 .await;
4530 recorded.borrow().clone()
4531 });
4532 h2_events(&mut client, &output)
4533 .into_iter()
4534 .filter_map(|event| match event {
4535 Event::WindowUpdateReceived {
4536 stream_id,
4537 increment,
4538 } => Some((stream_id, increment)),
4539 _ => None,
4540 })
4541 .collect()
4542 }
4543
4544 #[cfg(feature = "http2")]
4545 #[test]
4546 fn http2_stream_reset_drops_the_in_flight_handler() {
4547 use shiguredo_http2::ErrorCode;
4548
4549 let (mut client, stream, mut request) = h2_client("GET", "/forever", None);
4550 client
4551 .reset_stream(stream, ErrorCode::Cancel)
4552 .expect("the client resets its stream");
4553 while let Some(output) = client.poll_output() {
4554 request.extend_from_slice(&output);
4555 }
4556
4557 let server = Server::new(
4558 ExecutableApp::new([never_completing_operation()]).expect("the waiting app compiles"),
4559 );
4560 let runtime = Runtime::new().expect("the Compio runtime starts");
4561 let completed = runtime.block_on(async {
4562 let mut transport = ScriptedTransport::closing(&request);
4563 compio::time::timeout(
4564 Duration::from_secs(2),
4565 server.serve_http2_io(&mut transport),
4566 )
4567 .await
4568 .is_ok()
4569 });
4570
4571 assert!(
4572 completed,
4573 "a reset stream must drop its handler instead of holding the connection open"
4574 );
4575 }
4576
4577 #[test]
4578 fn generic_transport_streams_a_body_before_the_client_finishes_it() {
4579 clear_events();
4580 let server = Server::new(
4581 ExecutableApp::new([logging_upload_operation()]).expect("the upload app compiles"),
4582 );
4583 let mut transport = StagedTransport::new(vec![
4584 (
4585 "head",
4586 b"POST /staged-upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 8\r\nconnection: close\r\n\r\n".to_vec(),
4587 ),
4588 ("first", b"aaaa".to_vec()),
4589 ("last", b"bbbb".to_vec()),
4590 ]);
4591
4592 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4593
4594 let events = recorded_events();
4595 assert!(
4596 event_position(&events, "handler:aaaa") < event_position(&events, "client:last"),
4597 "the generic transport buffered the whole body before dispatch: {events:?}"
4598 );
4599 let written = transport.written();
4600 assert!(written.starts_with("HTTP/1.1 200 "), "{written}");
4601 assert!(written.ends_with("bytes=8"), "{written}");
4602 }
4603
4604 #[test]
4605 fn owned_generic_transport_streams_uploads_like_the_plaintext_socket() {
4606 let server = Server::new(
4607 ExecutableApp::new([upload_operation("/upload", "upload.consume", true)])
4608 .expect("the upload app compiles"),
4609 );
4610 let transport = ScriptedTransport::closing(
4611 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 11\r\nconnection: close\r\n\r\nhello world",
4612 );
4613 let written = transport.output();
4614
4615 future::block_on(server.serve_owned_io(transport)).expect("the connection completes");
4616
4617 let written = String::from_utf8_lossy(&written.borrow()).into_owned();
4618 assert!(written.starts_with("HTTP/1.1 200 "), "{written}");
4619 assert!(
4620 written.ends_with("bytes=11"),
4621 "the owned generic transport did not reach the streaming seam: {written}"
4622 );
4623 }
4624
4625 #[test]
4626 fn generic_transport_streams_a_chunked_upload() {
4627 let server = Server::new(
4628 ExecutableApp::new([upload_operation("/upload", "upload.consume", true)])
4629 .expect("the upload app compiles"),
4630 );
4631 let mut transport = ScriptedTransport::closing(
4632 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ntransfer-encoding: chunked\r\nconnection: close\r\n\r\n5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n",
4633 );
4634
4635 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4636
4637 let written = transport.written();
4638 assert!(written.ends_with("bytes=11"), "{written}");
4639 }
4640
4641 #[test]
4642 fn generic_transport_enforces_the_chunk_count_limit_while_streaming() {
4643 let limits = ServerLimits::new().with_max_chunks(1);
4644 let server = Server::new(
4645 ExecutableApp::new([upload_operation("/upload", "upload.consume", true)])
4646 .expect("the upload app compiles"),
4647 )
4648 .with_limits(limits);
4649 let mut transport = ScriptedTransport::closing(
4650 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ntransfer-encoding: chunked\r\nconnection: close\r\n\r\n5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n",
4651 );
4652
4653 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4654
4655 let written = transport.written();
4656 assert!(written.starts_with("HTTP/1.1 413 "), "{written}");
4657 assert!(written.contains("too_many_chunks"), "{written}");
4658 }
4659
4660 #[test]
4661 fn generic_transport_answers_a_truncated_upload_with_400() {
4662 let server = Server::new(
4663 ExecutableApp::new([upload_operation("/upload", "upload.consume", true)])
4664 .expect("the upload app compiles"),
4665 );
4666 let mut transport = ScriptedTransport::closing(
4667 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 11\r\nconnection: close\r\n\r\nhell",
4668 );
4669
4670 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4671
4672 let written = transport.written();
4673 assert!(written.starts_with("HTTP/1.1 400 "), "{written}");
4674 assert!(written.contains("incomplete_body"), "{written}");
4675 }
4676
4677 #[test]
4678 fn interim_continue_precedes_the_request_body_on_the_generic_transport() {
4679 clear_events();
4680 let server = Server::new(
4681 ExecutableApp::new([buffered_post_operation()]).expect("the buffered app compiles"),
4682 );
4683 let mut transport = StagedTransport::new(vec![
4684 (
4685 "head",
4686 b"POST /buffered HTTP/1.1\r\nhost: localhost\r\ncontent-length: 5\r\nexpect: 100-continue\r\nconnection: close\r\n\r\n".to_vec(),
4687 ),
4688 ("body", b"hello".to_vec()),
4689 ]);
4690
4691 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4692
4693 assert_eq!(
4694 recorded_events(),
4695 ["client:head", "server:continue", "client:body"],
4696 "the interim response must be written before the body is read"
4697 );
4698 let written = transport.written();
4699 assert!(
4700 written.starts_with("HTTP/1.1 100 Continue\r\n\r\nHTTP/1.1 200 "),
4701 "{written}"
4702 );
4703 assert_eq!(
4704 written.matches("100 Continue").count(),
4705 1,
4706 "the interim response was sent more than once: {written}"
4707 );
4708 }
4709
4710 #[test]
4711 fn interim_continue_precedes_a_streamed_upload_on_the_generic_transport() {
4712 clear_events();
4713 let server = Server::new(
4714 ExecutableApp::new([logging_upload_operation()]).expect("the upload app compiles"),
4715 );
4716 let mut transport = StagedTransport::new(vec![
4717 (
4718 "head",
4719 b"POST /staged-upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 4\r\nexpect: 100-continue\r\nconnection: close\r\n\r\n".to_vec(),
4720 ),
4721 ("body", b"aaaa".to_vec()),
4722 ]);
4723
4724 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4725
4726 assert_eq!(
4727 recorded_events(),
4728 [
4729 "client:head",
4730 "server:continue",
4731 "client:body",
4732 "handler:aaaa"
4733 ],
4734 "the streaming path must answer the expectation before reading the body"
4735 );
4736 }
4737
4738 #[test]
4739 fn interim_continue_is_written_on_the_plaintext_socket() {
4740 let response = native_exchange(
4741 vec![buffered_post_operation()],
4742 b"POST /buffered HTTP/1.1\r\nhost: localhost\r\ncontent-length: 5\r\nexpect: 100-continue\r\nconnection: close\r\n\r\nhello"
4743 .to_vec(),
4744 );
4745
4746 let response = String::from_utf8_lossy(&response).into_owned();
4747 assert!(
4748 response.starts_with("HTTP/1.1 100 Continue\r\n\r\nHTTP/1.1 200 "),
4749 "{response}"
4750 );
4751 assert_eq!(
4752 response.matches("100 Continue").count(),
4753 1,
4754 "the interim response was sent more than once: {response}"
4755 );
4756 }
4757
4758 #[test]
4759 fn unknown_expectation_is_answered_with_417_on_the_generic_transport() {
4760 let server = Server::new(
4761 ExecutableApp::new([buffered_post_operation()]).expect("the buffered app compiles"),
4762 );
4763 let mut transport = ScriptedTransport::closing(
4764 b"POST /buffered HTTP/1.1\r\nhost: localhost\r\ncontent-length: 5\r\nexpect: the-moon\r\n\r\nhello",
4765 );
4766
4767 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4768
4769 let written = transport.written();
4770 assert!(written.starts_with("HTTP/1.1 417 "), "{written}");
4771 assert!(written.contains("expectation_failed"), "{written}");
4772 assert!(!written.contains("100 Continue"), "{written}");
4773 }
4774
4775 #[test]
4776 fn unknown_expectation_is_answered_with_417_on_the_plaintext_socket() {
4777 let response = native_exchange(
4778 vec![buffered_post_operation()],
4779 b"POST /buffered HTTP/1.1\r\nhost: localhost\r\ncontent-length: 5\r\nexpect: the-moon\r\n\r\nhello"
4780 .to_vec(),
4781 );
4782
4783 let response = String::from_utf8_lossy(&response).into_owned();
4784 assert!(response.starts_with("HTTP/1.1 417 "), "{response}");
4785 assert!(response.contains("expectation_failed"), "{response}");
4786 }
4787
4788 #[test]
4789 fn an_oversized_expected_body_is_answered_with_413_without_an_interim_response() {
4790 let server = Server::new(
4791 ExecutableApp::new([buffered_post_operation()]).expect("the buffered app compiles"),
4792 )
4793 .with_limits(ServerLimits::new().with_max_body_bytes(8));
4794 let mut transport = ScriptedTransport::closing(
4795 b"POST /buffered HTTP/1.1\r\nhost: localhost\r\ncontent-length: 4096\r\nexpect: 100-continue\r\n\r\n",
4796 );
4797
4798 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4799
4800 let written = transport.written();
4801 assert!(written.starts_with("HTTP/1.1 413 "), "{written}");
4802 assert!(written.contains("payload_too_large"), "{written}");
4803 assert!(
4804 !written.contains("100 Continue"),
4805 "a rejected request must not be invited to send its body: {written}"
4806 );
4807 }
4808
4809 #[test]
4810 fn an_oversized_streamed_body_is_answered_with_413_on_the_plaintext_socket() {
4811 let response = native_exchange_with_limits(
4812 vec![upload_operation("/upload", "upload.consume", true)],
4813 b"POST /upload HTTP/1.1\r\nhost: localhost\r\ncontent-length: 4096\r\nexpect: 100-continue\r\n\r\n"
4814 .to_vec(),
4815 ServerLimits::new().with_max_body_bytes(8),
4816 );
4817
4818 let response = String::from_utf8_lossy(&response).into_owned();
4819 assert!(
4820 !response.contains("100 Continue"),
4821 "a rejected upload must not be invited to send its body: {response}"
4822 );
4823 assert!(response.contains("payload_too_large"), "{response}");
4824 }
4825
4826 #[test]
4827 fn an_http_1_0_expectation_is_ignored() {
4828 let server = Server::new(
4829 ExecutableApp::new([buffered_post_operation()]).expect("the buffered app compiles"),
4830 );
4831 let mut transport = ScriptedTransport::closing(
4832 b"POST /buffered HTTP/1.0\r\nhost: localhost\r\ncontent-length: 5\r\nexpect: 100-continue\r\n\r\nhello",
4833 );
4834
4835 future::block_on(server.serve_io(&mut transport)).expect("the connection completes");
4836
4837 let written = transport.written();
4838 assert!(written.starts_with("HTTP/1.1 200 "), "{written}");
4839 assert!(
4840 !written.contains("100 Continue"),
4841 "an HTTP/1.0 peer must never receive an interim response: {written}"
4842 );
4843 }
4844
4845 #[test]
4846 fn worker_middleware_factory_is_applied_to_the_worker_app() {
4847 let config = super::WorkerConfig {
4848 factory: ping_app,
4849 max_body_bytes: super::DEFAULT_MAX_BODY_BYTES,
4850 openapi: None,
4851 middleware: Some(Arc::new(|| {
4852 vec![Rc::new(StampMiddleware) as Rc<dyn HttpMiddleware>]
4853 })),
4854 };
4855 let server_id = super::NEXT_MULTICORE_SERVER_ID.fetch_add(1, Ordering::Relaxed);
4856 let (app, created) = super::worker_app(server_id, &config);
4857 assert!(created);
4858
4859 let response = future::block_on(app.call(Request::new(HttpMethod::Get, "/ping")));
4860 super::take_worker_app(server_id);
4861
4862 assert_eq!(
4863 response.get_header("x-native-middleware"),
4864 Some("applied"),
4865 "the worker factory middleware did not run"
4866 );
4867 }
4868}