1use std::time::Duration;
4
5use rama_core::{
6 Service,
7 bytes::BytesMut,
8 error::{BoxError, ErrorContext},
9 io::{
10 PeekIoProvider, PrefixedIo, ReplayReader,
11 peek::{PeekTimeoutError, PeekTimeoutPolicy},
12 },
13 service::RejectService,
14 telemetry::tracing,
15};
16use rama_utils::octets::kib;
17use tokio::{io::AsyncReadExt as _, time::Instant};
18
19use crate::{
20 byte_sets::{is_control_byte, is_http_token_byte, is_scheme_first_byte, is_scheme_rest_byte},
21 uri::parser::validate_http_request_target,
22};
23
24pub const DEFAULT_HTTP1_REQUEST_LINE_MAX_SIZE: usize = kib(8);
30
31pub const DEFAULT_HTTP_PEEK_READ_BUFFER_SIZE: usize = 512;
33
34pub const KNOWN_NON_HTTP_PROTOCOL_METHODS: &[&str] =
41 &["PING", "EHLO", "HELO", "USER", "NICK", "SSH", "PROXY"];
42
43const INITIAL_HTTP_PEEK_READ_BUFFER_SIZE: usize = 128;
48
49#[non_exhaustive]
51#[derive(Debug, Clone, Copy, PartialEq, Eq)]
52pub struct HttpPeekConfig {
53 pub timeout: Option<Duration>,
55 pub timeout_policy: PeekTimeoutPolicy,
58 pub max_http1_request_line_size: usize,
60 pub read_buffer_size: usize,
62 pub skipped_http1_methods: &'static [&'static str],
67}
68
69impl Default for HttpPeekConfig {
70 fn default() -> Self {
71 Self {
72 timeout: None,
73 timeout_policy: PeekTimeoutPolicy::default(),
74 max_http1_request_line_size: DEFAULT_HTTP1_REQUEST_LINE_MAX_SIZE,
75 read_buffer_size: DEFAULT_HTTP_PEEK_READ_BUFFER_SIZE,
76 skipped_http1_methods: &[],
77 }
78 }
79}
80
81impl HttpPeekConfig {
82 #[must_use]
84 pub fn new() -> Self {
85 Self::default()
86 }
87
88 rama_utils::macros::generate_set_and_with! {
89 pub fn timeout_policy(mut self, timeout_policy: PeekTimeoutPolicy) -> Self {
94 self.timeout_policy = timeout_policy;
95 self
96 }
97 }
98
99 rama_utils::macros::generate_set_and_with! {
100 pub fn known_non_http_protocol_methods(mut self) -> Self {
102 self.skipped_http1_methods = KNOWN_NON_HTTP_PROTOCOL_METHODS;
103 self
104 }
105 }
106
107 rama_utils::macros::generate_set_and_with! {
108 pub fn skipped_http1_methods(
110 mut self,
111 skipped_http1_methods: &'static [&'static str],
112 ) -> Self {
113 self.skipped_http1_methods = skipped_http1_methods;
114 self
115 }
116 }
117}
118
119#[derive(Debug, Clone)]
128pub struct HttpPeekRouter<T, F = RejectService<(), NoHttpRejectError>> {
129 http_acceptor: T,
130 fallback: F,
131 peek_config: HttpPeekConfig,
132}
133
134#[derive(Debug, Clone)]
137pub struct HttpDualAcceptor<T, U> {
138 http1: T,
139 h2: U,
140}
141
142#[derive(Debug, Clone)]
145pub struct HttpAutoAcceptor<T>(T);
146
147#[derive(Debug, Clone)]
150pub struct Http1Acceptor<T>(T);
151
152#[derive(Debug, Clone)]
155pub struct H2Acceptor<T>(T);
156
157rama_utils::macros::error::static_str_error! {
158 #[doc = "non-http connection is rejected"]
159 pub struct NoHttpRejectError;
160}
161
162impl<T> HttpPeekRouter<HttpAutoAcceptor<T>> {
163 pub fn new(auto_acceptor: T) -> Self {
166 Self {
167 http_acceptor: HttpAutoAcceptor(auto_acceptor),
168 fallback: RejectService::new(NoHttpRejectError),
169 peek_config: HttpPeekConfig::default(),
170 }
171 }
172}
173
174impl<T> HttpPeekRouter<Http1Acceptor<T>> {
175 pub fn new_http1(http1_acceptor: T) -> Self {
178 Self {
179 http_acceptor: Http1Acceptor(http1_acceptor),
180 fallback: RejectService::new(NoHttpRejectError),
181 peek_config: HttpPeekConfig::default(),
182 }
183 }
184}
185
186impl<T> HttpPeekRouter<H2Acceptor<T>> {
187 pub fn new_h2(h2_acceptor: T) -> Self {
190 Self {
191 http_acceptor: H2Acceptor(h2_acceptor),
192 fallback: RejectService::new(NoHttpRejectError),
193 peek_config: HttpPeekConfig::default(),
194 }
195 }
196}
197
198impl<T> HttpPeekRouter<T> {
199 pub fn with_fallback<F>(self, fallback: F) -> HttpPeekRouter<T, F> {
201 HttpPeekRouter {
202 http_acceptor: self.http_acceptor,
203 fallback,
204 peek_config: self.peek_config,
205 }
206 }
207}
208
209impl<T, F> HttpPeekRouter<T, F> {
210 rama_utils::macros::generate_set_and_with! {
211 pub fn known_non_http_protocol_methods(mut self) -> Self {
213 self.peek_config = self.peek_config.with_known_non_http_protocol_methods();
214 self
215 }
216 }
217
218 rama_utils::macros::generate_set_and_with! {
219 pub fn peek_timeout(mut self, peek_timeout: Option<Duration>) -> Self {
225 self.peek_config.timeout = peek_timeout;
226 self
227 }
228 }
229
230 rama_utils::macros::generate_set_and_with! {
231 pub fn peek_timeout_policy(mut self, peek_timeout_policy: PeekTimeoutPolicy) -> Self {
236 self.peek_config.timeout_policy = peek_timeout_policy;
237 self
238 }
239 }
240
241 rama_utils::macros::generate_set_and_with! {
242 pub fn peek_config(mut self, peek_config: HttpPeekConfig) -> Self {
244 self.peek_config = peek_config;
245 self
246 }
247 }
248
249 rama_utils::macros::generate_set_and_with! {
250 pub fn skipped_http1_methods(
252 mut self,
253 skipped_http1_methods: &'static [&'static str],
254 ) -> Self {
255 self.peek_config.skipped_http1_methods = skipped_http1_methods;
256 self
257 }
258 }
259}
260
261impl<T, U> HttpPeekRouter<HttpDualAcceptor<T, U>> {
262 pub fn new_dual(http1_acceptor: T, h2_acceptor: U) -> Self {
265 Self {
266 http_acceptor: HttpDualAcceptor {
267 http1: http1_acceptor,
268 h2: h2_acceptor,
269 },
270 fallback: RejectService::new(NoHttpRejectError),
271 peek_config: HttpPeekConfig::default(),
272 }
273 }
274}
275
276impl<PeekableInput, Output, T, F> Service<PeekableInput> for HttpPeekRouter<HttpAutoAcceptor<T>, F>
277where
278 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
279 Output: Send + 'static,
280 T: Service<
281 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
282 Output = Output,
283 Error: Into<BoxError>,
284 >,
285 F: Service<
286 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
287 Output = Output,
288 Error: Into<BoxError>,
289 >,
290{
291 type Output = Output;
292 type Error = BoxError;
293
294 async fn serve(&self, input: PeekableInput) -> Result<Self::Output, Self::Error> {
295 let (version, peek_input) = peek_http_input_with_config(input, self.peek_config).await?;
296 if version.is_some() {
297 tracing::debug!(
298 "http peek [auto]: HTTP detect: version = {version:?}; continue with http_acceptor svc"
299 );
300 self.http_acceptor
301 .0
302 .serve(peek_input)
303 .await
304 .into_box_error()
305 } else {
306 tracing::debug!("http peek [auto]: HTTP not detect: continue with fallback svc");
307 self.fallback.serve(peek_input).await.into_box_error()
308 }
309 }
310}
311
312impl<PeekableInput, Output, T, F> Service<PeekableInput> for HttpPeekRouter<Http1Acceptor<T>, F>
313where
314 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
315 Output: Send + 'static,
316 T: Service<
317 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
318 Output = Output,
319 Error: Into<BoxError>,
320 >,
321 F: Service<
322 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
323 Output = Output,
324 Error: Into<BoxError>,
325 >,
326{
327 type Output = Output;
328 type Error = BoxError;
329
330 async fn serve(&self, input: PeekableInput) -> Result<Self::Output, Self::Error> {
331 let (version, peek_input) = peek_http_input_with_config(input, self.peek_config).await?;
332 if version == Some(HttpPeekVersion::Http1x) {
333 tracing::debug!("http peek: serve[http1]: http/1x acceptor; version = {version:?}");
334 self.http_acceptor
335 .0
336 .serve(peek_input)
337 .await
338 .into_box_error()
339 } else {
340 tracing::debug!("http peek: serve[http1]: fallback; version = {version:?}");
341 self.fallback.serve(peek_input).await.into_box_error()
342 }
343 }
344}
345
346impl<PeekableInput, Output, T, F> Service<PeekableInput> for HttpPeekRouter<H2Acceptor<T>, F>
347where
348 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
349 Output: Send + 'static,
350 T: Service<
351 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
352 Output = Output,
353 Error: Into<BoxError>,
354 >,
355 F: Service<
356 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
357 Output = Output,
358 Error: Into<BoxError>,
359 >,
360{
361 type Output = Output;
362 type Error = BoxError;
363
364 async fn serve(&self, input: PeekableInput) -> Result<Self::Output, Self::Error> {
365 let (version, peek_input) = peek_http_input_with_config(input, self.peek_config).await?;
366 if version == Some(HttpPeekVersion::H2) {
367 tracing::debug!("http peek: serve[h2]: http acceptor; version = {version:?}");
368 self.http_acceptor
369 .0
370 .serve(peek_input)
371 .await
372 .into_box_error()
373 } else {
374 tracing::debug!("http peek: serve[h2]: fallback; version = {version:?}");
375 self.fallback.serve(peek_input).await.into_box_error()
376 }
377 }
378}
379
380impl<PeekableInput, Output, T, U, F> Service<PeekableInput>
381 for HttpPeekRouter<HttpDualAcceptor<T, U>, F>
382where
383 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
384 Output: Send + 'static,
385 T: Service<
386 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
387 Output = Output,
388 Error: Into<BoxError>,
389 >,
390 U: Service<
391 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
392 Output = Output,
393 Error: Into<BoxError>,
394 >,
395 F: Service<
396 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
397 Output = Output,
398 Error: Into<BoxError>,
399 >,
400{
401 type Output = Output;
402 type Error = BoxError;
403
404 async fn serve(&self, input: PeekableInput) -> Result<Self::Output, Self::Error> {
405 let (version, peek_input) = peek_http_input_with_config(input, self.peek_config).await?;
406 match version {
407 Some(HttpPeekVersion::H2) => {
408 tracing::trace!("http peek: serve[dual]: h2 acceptor; version = {version:?}");
409 self.http_acceptor
410 .h2
411 .serve(peek_input)
412 .await
413 .into_box_error()
414 }
415 Some(HttpPeekVersion::Http1x) => {
416 tracing::trace!("http peek: serve[dual]: http/1x acceptor; version = {version:?}");
417 self.http_acceptor
418 .http1
419 .serve(peek_input)
420 .await
421 .into_box_error()
422 }
423 None => {
424 tracing::trace!("http peek: serve[dual]: fallback; version = {version:?}");
425 self.fallback.serve(peek_input).await.into_box_error()
426 }
427 }
428 }
429}
430
431#[derive(Debug, Clone, Copy, PartialEq, Eq)]
432pub enum HttpPeekVersion {
433 Http1x,
434 H2,
435}
436
437#[derive(Debug, Clone, Copy)]
438enum Http1PeekState {
439 LeadingLf,
441 Method {
442 start: usize,
443 len: usize,
444 },
445 Target {
446 method_start: usize,
447 method_end: usize,
448 target: HttpRequestTargetState,
449 },
450 Http09Lf {
451 method_start: usize,
452 method_end: usize,
453 target_end: usize,
454 },
455 Version {
456 method_start: usize,
457 method_end: usize,
458 target_end: usize,
459 offset: usize,
460 minor: u8,
461 },
462 Matched,
463 Invalid,
464}
465
466#[derive(Debug)]
467struct HttpPeekState {
468 http1: Http1PeekState,
469 h2_offset: Option<usize>,
470 max_http1_request_line_size: usize,
471 skipped_http1_methods: &'static [&'static str],
472}
473
474#[derive(Debug, Clone, Copy, PartialEq, Eq)]
475enum HttpPeekDecision {
476 Continue,
477 Matched(HttpPeekVersion),
478 Reject,
479}
480
481#[derive(Debug, Clone, Copy)]
482enum HttpRequestTargetState {
483 Start,
484 Origin,
485 Asterisk,
486 Scheme { len: usize },
487 Absolute,
488 Authority,
489}
490
491impl HttpRequestTargetState {
492 fn push(self, byte: u8) -> Option<Self> {
493 if !is_http_request_target_prefix_byte(byte) {
494 return None;
495 }
496
497 match self {
498 Self::Start if byte == b'/' => Some(Self::Origin),
499 Self::Start if byte == b'*' => Some(Self::Asterisk),
500 Self::Start if is_scheme_first_byte(byte) => Some(Self::Scheme { len: 1 }),
501 Self::Origin if byte != b'#' => Some(Self::Origin),
502 Self::Scheme { len: _ } if byte == b':' => Some(Self::Absolute),
503 Self::Scheme { len }
504 if len < crate::proto::MAX_SCHEME_LEN && is_scheme_rest_byte(byte) =>
505 {
506 Some(Self::Scheme { len: len + 1 })
507 }
508 Self::Absolute if byte != b'#' => Some(Self::Absolute),
509 Self::Authority if !matches!(byte, b'/' | b'?' | b'#') => Some(Self::Authority),
510 _ => None,
511 }
512 }
513}
514
515impl HttpPeekState {
516 fn new(
517 max_http1_request_line_size: usize,
518 skipped_http1_methods: &'static [&'static str],
519 ) -> Self {
520 Self {
521 http1: if max_http1_request_line_size == 0 {
522 Http1PeekState::Invalid
523 } else {
524 Http1PeekState::Method { start: 0, len: 0 }
525 },
526 h2_offset: Some(0),
527 max_http1_request_line_size,
528 skipped_http1_methods,
529 }
530 }
531
532 fn max_peek_len(&self) -> usize {
533 let http1 = if matches!(
534 self.http1,
535 Http1PeekState::Invalid | Http1PeekState::Matched
536 ) {
537 0
538 } else {
539 self.max_http1_request_line_size
540 };
541 let h2 = self.h2_offset.map(|_| H2_MAGIC_PREFIX.len()).unwrap_or(0);
542 http1.max(h2)
543 }
544
545 fn push_byte(&mut self, byte: u8, total_len: usize, buffer: &[u8]) -> HttpPeekDecision {
546 if let Some(offset) = self.h2_offset {
547 if H2_MAGIC_PREFIX.get(offset) == Some(&byte) {
548 let next = offset + 1;
549 if next == H2_MAGIC_PREFIX.len() {
550 tracing::trace!(version = "HTTP/2", "HTTP peek matched client preface");
551 return HttpPeekDecision::Matched(HttpPeekVersion::H2);
552 }
553 self.h2_offset = Some(next);
554 } else {
555 self.h2_offset = None;
556 }
557 }
558
559 self.push_http1_byte(byte, total_len, buffer);
560
561 if matches!(self.http1, Http1PeekState::Matched) {
562 return HttpPeekDecision::Matched(HttpPeekVersion::Http1x);
563 }
564
565 if total_len >= self.max_http1_request_line_size {
566 self.http1 = Http1PeekState::Invalid;
567 }
568
569 if matches!(self.http1, Http1PeekState::Invalid) && self.h2_offset.is_none() {
570 HttpPeekDecision::Reject
571 } else {
572 HttpPeekDecision::Continue
573 }
574 }
575
576 fn push_http1_byte(&mut self, byte: u8, total_len: usize, buffer: &[u8]) {
577 const VERSION_PREFIX: &[u8] = b"HTTP/1.";
579
580 let state = core::mem::replace(&mut self.http1, Http1PeekState::Invalid);
581 self.http1 = match state {
582 Http1PeekState::Method { start, len } => {
583 if is_http_token_byte(byte) {
584 let len = len + 1;
585 let method = &buffer[start..start + len];
586 if self
587 .skipped_http1_methods
588 .iter()
589 .any(|skipped| method == skipped.as_bytes())
590 {
591 tracing::trace!(
592 method = ?core::str::from_utf8(method).ok(),
593 "HTTP/1 peek rejected configured method"
594 );
595 Http1PeekState::Invalid
596 } else {
597 Http1PeekState::Method { start, len }
598 }
599 } else if byte == b' ' && len > 0 {
600 Http1PeekState::Target {
601 method_start: start,
602 method_end: start + len,
603 target: if &buffer[start..start + len] == b"CONNECT" {
604 HttpRequestTargetState::Authority
605 } else {
606 HttpRequestTargetState::Start
607 },
608 }
609 } else if len == 0 && byte == b'\n' {
610 Http1PeekState::Method {
612 start: total_len,
613 len: 0,
614 }
615 } else if len == 0 && byte == b'\r' {
616 Http1PeekState::LeadingLf
617 } else {
618 tracing::trace!(byte, "HTTP/1 peek rejected invalid method byte");
619 Http1PeekState::Invalid
620 }
621 }
622 Http1PeekState::LeadingLf => {
623 if byte == b'\n' {
624 Http1PeekState::Method {
625 start: total_len,
626 len: 0,
627 }
628 } else {
629 tracing::trace!(byte, "HTTP/1 peek rejected bare CR before request-line");
630 Http1PeekState::Invalid
631 }
632 }
633 Http1PeekState::Target {
634 method_start,
635 method_end,
636 target,
637 } => {
638 if matches!(byte, b' ' | b'\r' | b'\n') {
639 let target_start = method_end + 1;
640 let target_end = total_len - 1;
641 let method = &buffer[method_start..method_end];
642 let target = &buffer[target_start..target_end];
643 let authority_form = method == b"CONNECT";
644 if target == b"*" && method != b"OPTIONS" {
645 tracing::trace!(
646 "HTTP/1 peek rejected asterisk-form for method other than OPTIONS"
647 );
648 Http1PeekState::Invalid
649 } else {
650 match validate_http_request_target(target, authority_form) {
651 Ok(()) if byte == b' ' => Http1PeekState::Version {
652 method_start,
653 method_end,
654 target_end,
655 offset: 0,
656 minor: 0,
657 },
658 Ok(()) if method == b"GET" => {
659 if byte == b'\n' {
660 trace_http1_match(
662 buffer,
663 method_start,
664 method_end,
665 target_end,
666 "HTTP/0.9",
667 );
668 Http1PeekState::Matched
669 } else {
670 Http1PeekState::Http09Lf {
671 method_start,
672 method_end,
673 target_end,
674 }
675 }
676 }
677 Ok(()) => {
678 tracing::trace!("HTTP/0.9 peek rejected method other than GET");
679 Http1PeekState::Invalid
680 }
681 Err(err) => {
682 tracing::trace!(%err, "HTTP/1 peek rejected invalid request target");
683 Http1PeekState::Invalid
684 }
685 }
686 }
687 } else if let Some(target) = target.push(byte) {
688 Http1PeekState::Target {
689 method_start,
690 method_end,
691 target,
692 }
693 } else {
694 tracing::trace!(byte, "HTTP/1 peek rejected invalid request-target byte");
695 Http1PeekState::Invalid
696 }
697 }
698 Http1PeekState::Http09Lf {
699 method_start,
700 method_end,
701 target_end,
702 } => {
703 if byte == b'\n' {
704 trace_http1_match(buffer, method_start, method_end, target_end, "HTTP/0.9");
705 Http1PeekState::Matched
706 } else {
707 tracing::trace!(byte, "HTTP/0.9 peek rejected invalid line ending");
708 Http1PeekState::Invalid
709 }
710 }
711 Http1PeekState::Version {
712 method_start,
713 method_end,
714 target_end,
715 offset,
716 minor,
717 } => {
718 if offset < VERSION_PREFIX.len() {
719 if byte == VERSION_PREFIX[offset] {
720 Http1PeekState::Version {
721 method_start,
722 method_end,
723 target_end,
724 offset: offset + 1,
725 minor,
726 }
727 } else {
728 tracing::trace!(byte, offset, "HTTP/1 peek rejected invalid version byte");
729 Http1PeekState::Invalid
730 }
731 } else if offset == VERSION_PREFIX.len() {
732 if byte.is_ascii_digit() {
733 Http1PeekState::Version {
734 method_start,
735 method_end,
736 target_end,
737 offset: offset + 1,
738 minor: byte,
739 }
740 } else {
741 tracing::trace!(byte, offset, "HTTP/1 peek rejected invalid version byte");
742 Http1PeekState::Invalid
743 }
744 } else if byte == b'\n' {
745 trace_http1_version_match(buffer, method_start, method_end, target_end, minor);
747 Http1PeekState::Matched
748 } else if byte == b'\r' && offset == VERSION_PREFIX.len() + 1 {
749 Http1PeekState::Version {
750 method_start,
751 method_end,
752 target_end,
753 offset: offset + 1,
754 minor,
755 }
756 } else {
757 tracing::trace!(byte, offset, "HTTP/1 peek rejected invalid line ending");
758 Http1PeekState::Invalid
759 }
760 }
761 state @ (Http1PeekState::Matched | Http1PeekState::Invalid) => state,
762 };
763 }
764}
765
766#[inline]
767fn trace_http1_match(
768 buffer: &[u8],
769 method_start: usize,
770 method_end: usize,
771 target_end: usize,
772 version: &'static str,
773) {
774 let method = unsafe { core::str::from_utf8_unchecked(&buffer[method_start..method_end]) };
777 let target = unsafe { core::str::from_utf8_unchecked(&buffer[method_end + 1..target_end]) };
778 tracing::trace!(method, target, version, "HTTP/1 peek matched request-line");
779}
780
781#[inline]
782fn trace_http1_version_match(
783 buffer: &[u8],
784 method_start: usize,
785 method_end: usize,
786 target_end: usize,
787 minor: u8,
788) {
789 const VERSIONS: [&str; 10] = [
790 "HTTP/1.0", "HTTP/1.1", "HTTP/1.2", "HTTP/1.3", "HTTP/1.4", "HTTP/1.5", "HTTP/1.6",
791 "HTTP/1.7", "HTTP/1.8", "HTTP/1.9",
792 ];
793 let version = VERSIONS[usize::from(minor.saturating_sub(b'0')).min(9)];
794 trace_http1_match(buffer, method_start, method_end, target_end, version);
795}
796
797#[inline]
798fn is_http_request_target_prefix_byte(byte: u8) -> bool {
799 byte != b' ' && !is_control_byte(byte)
800}
801
802pub async fn peek_http_input<PeekableInput>(
808 input: PeekableInput,
809 timeout: Option<Duration>,
810) -> Result<
811 (
812 Option<HttpPeekVersion>,
813 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
814 ),
815 BoxError,
816>
817where
818 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
819{
820 peek_http_input_with_timeout_policy(input, timeout, PeekTimeoutPolicy::FailOpen).await
821}
822
823pub async fn peek_http_input_with_timeout_policy<PeekableInput>(
830 input: PeekableInput,
831 timeout: Option<Duration>,
832 timeout_policy: PeekTimeoutPolicy,
833) -> Result<
834 (
835 Option<HttpPeekVersion>,
836 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
837 ),
838 BoxError,
839>
840where
841 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
842{
843 peek_http_input_with_config(
844 input,
845 HttpPeekConfig {
846 timeout,
847 timeout_policy,
848 ..Default::default()
849 },
850 )
851 .await
852}
853
854pub async fn peek_http_input_with_config<PeekableInput>(
867 mut input: PeekableInput,
868 config: HttpPeekConfig,
869) -> Result<
870 (
871 Option<HttpPeekVersion>,
872 PeekableInput::Mapped<HttpPrefixedIo<PeekableInput::PeekIo>>,
873 ),
874 BoxError,
875>
876where
877 PeekableInput: PeekIoProvider<PeekIo: Unpin>,
878{
879 let mut state = HttpPeekState::new(
880 config.max_http1_request_line_size,
881 config.skipped_http1_methods,
882 );
883 let max_peek_len = state.max_peek_len();
884 let read_buffer_size = config.read_buffer_size.max(1);
885 let mut next_read_buffer_size = read_buffer_size
886 .min(INITIAL_HTTP_PEEK_READ_BUFFER_SIZE)
887 .min(max_peek_len);
888 let mut buffer = BytesMut::with_capacity(next_read_buffer_size);
889 let mut total_len = 0usize;
890 let mut matched_version = None;
891 let deadline = config.timeout.map(|duration| Instant::now() + duration);
892
893 'peek: loop {
894 let remaining = state.max_peek_len().saturating_sub(total_len);
895 if remaining == 0 {
896 break;
897 }
898
899 let read_capacity = remaining.min(next_read_buffer_size);
900 buffer.reserve(read_capacity);
901 let read_start = buffer.len();
902 let mut limited = input.peek_io_mut().take(read_capacity as u64);
903 let read = limited.read_buf(&mut buffer);
904 let read_size = match deadline {
905 Some(deadline) => match tokio::time::timeout_at(deadline, read).await {
906 Ok(Ok(size)) => size,
907 Ok(Err(err)) => {
908 tracing::debug!(%err, "HTTP peek read failed");
909 break;
910 }
911 Err(err) => {
912 tracing::debug!(%err, "HTTP peek timed out");
913 if config.timeout_policy == PeekTimeoutPolicy::FailClosed {
914 return Err(PeekTimeoutError::new().into());
915 }
916 break;
917 }
918 },
919 None => match read.await {
920 Ok(size) => size,
921 Err(err) => {
922 tracing::debug!(%err, "HTTP peek read failed");
923 break;
924 }
925 },
926 };
927
928 let Some(_) = core::num::NonZeroUsize::new(read_size) else {
929 break;
930 };
931
932 if read_size == read_capacity {
933 next_read_buffer_size = next_read_buffer_size
934 .saturating_mul(2)
935 .min(read_buffer_size);
936 }
937
938 for index in read_start..buffer.len() {
939 total_len = index + 1;
940 match state.push_byte(buffer[index], total_len, &buffer) {
941 HttpPeekDecision::Continue => {}
942 HttpPeekDecision::Matched(version) => {
943 matched_version = Some(version);
944 break 'peek;
945 }
946 HttpPeekDecision::Reject => {
947 break 'peek;
948 }
949 }
950 }
951 total_len = buffer.len();
952 }
953
954 tracing::trace!(
955 version = ?matched_version,
956 peek_size = buffer.len(),
957 "HTTP peek read loop finished"
958 );
959
960 let peek = ReplayReader::new(buffer.freeze());
961 let peek_input = input.map_peek_io(|io| PrefixedIo::new(peek, io));
962
963 Ok((matched_version, peek_input))
964}
965
966const H2_MAGIC_PREFIX: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
967
968pub type HttpPrefixedIo<S> = PrefixedIo<ReplayReader, S>;
970
971#[cfg(test)]
972mod test {
973 use core::convert::Infallible;
974 use std::{
975 io,
976 pin::Pin,
977 sync::{
978 Arc,
979 atomic::{AtomicUsize, Ordering},
980 },
981 task::{Context, Poll},
982 };
983
984 use super::*;
985
986 use parking_lot::Mutex;
987 use rama_core::io::Io;
988 use rama_core::{
989 ServiceInput,
990 bytes::Bytes,
991 futures::{StreamExt as _, async_stream::stream_fn},
992 service::{RejectError, service_fn},
993 stream::io::StreamReader,
994 };
995 use tokio::io::{AsyncRead, ReadBuf};
996
997 async fn peek_bytes(
998 content: &[u8],
999 config: HttpPeekConfig,
1000 ) -> (Option<HttpPeekVersion>, Vec<u8>) {
1001 let input = ServiceInput::new(std::io::Cursor::new(content.to_vec()));
1002 let (version, mut input) = peek_http_input_with_config(input, config).await.unwrap();
1003 let mut replayed = Vec::new();
1004 input.read_to_end(&mut replayed).await.unwrap();
1005 (version, replayed)
1006 }
1007
1008 async fn peek_fragmented_bytes(content: &'static [u8]) -> (Option<HttpPeekVersion>, Vec<u8>) {
1009 let reader = StreamReader::new(rama_core::futures::stream::iter(
1010 content
1011 .iter()
1012 .map(|&byte| Ok::<_, std::io::Error>(Bytes::copy_from_slice(&[byte]))),
1013 ));
1014 let io = Box::pin(tokio::io::join(reader, tokio::io::sink()));
1015 let (version, mut io) = peek_http_input(io, None).await.unwrap();
1016 let mut replayed = Vec::new();
1017 io.read_to_end(&mut replayed).await.unwrap();
1018 (version, replayed)
1019 }
1020
1021 fn state_decision(content: &[u8]) -> HttpPeekDecision {
1022 let mut state = HttpPeekState::new(DEFAULT_HTTP1_REQUEST_LINE_MAX_SIZE, &[]);
1023 let mut decision = HttpPeekDecision::Continue;
1024 for (index, &byte) in content.iter().enumerate() {
1025 decision = state.push_byte(byte, index + 1, &content[..=index]);
1026 if decision != HttpPeekDecision::Continue {
1027 break;
1028 }
1029 }
1030 decision
1031 }
1032
1033 fn state_decision_with_skipped_methods(
1034 content: &[u8],
1035 skipped_http1_methods: &'static [&'static str],
1036 ) -> HttpPeekDecision {
1037 let mut state =
1038 HttpPeekState::new(DEFAULT_HTTP1_REQUEST_LINE_MAX_SIZE, skipped_http1_methods);
1039 let mut decision = HttpPeekDecision::Continue;
1040 for (index, &byte) in content.iter().enumerate() {
1041 decision = state.push_byte(byte, index + 1, &content[..=index]);
1042 if decision != HttpPeekDecision::Continue {
1043 break;
1044 }
1045 }
1046 decision
1047 }
1048
1049 #[test]
1050 fn test_http1_state_rejects_at_first_impossible_byte() {
1051 assert_eq!(HttpPeekDecision::Reject, state_decision(b" "));
1052 assert_eq!(HttpPeekDecision::Reject, state_decision(b"\rG"));
1053 assert_eq!(HttpPeekDecision::Reject, state_decision(b"\r\r"));
1054 assert_eq!(HttpPeekDecision::Continue, state_decision(b"\r"));
1055 assert_eq!(HttpPeekDecision::Continue, state_decision(b"\r\n\nGET"));
1056 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GE("));
1057 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET \t"));
1058 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET \x7f"));
1059 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET /\t"));
1060 assert_eq!(HttpPeekDecision::Reject, state_decision(b"CONNECT h\t"));
1061 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET ["));
1062 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET *x"));
1063 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET ht!"));
1064 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET /#"));
1065 assert_eq!(HttpPeekDecision::Reject, state_decision(b"GET http://x#"));
1066 assert_eq!(HttpPeekDecision::Reject, state_decision(b"CONNECT host/"));
1067
1068 assert_eq!(HttpPeekDecision::Continue, state_decision(b"GET /\xc3"));
1071 assert_eq!(HttpPeekDecision::Continue, state_decision(b"GET http:x"));
1072 assert_eq!(HttpPeekDecision::Continue, state_decision(b"GET http:/x"));
1073 assert_eq!(HttpPeekDecision::Continue, state_decision(b"CONNECT ["));
1074
1075 let mut max_scheme = b"GET ".to_vec();
1076 max_scheme.extend(core::iter::repeat_n(b'a', crate::proto::MAX_SCHEME_LEN));
1077 assert_eq!(HttpPeekDecision::Continue, state_decision(&max_scheme));
1078 max_scheme.push(b':');
1079 assert_eq!(HttpPeekDecision::Continue, state_decision(&max_scheme));
1080
1081 let mut oversized_scheme = b"GET ".to_vec();
1082 oversized_scheme.extend(core::iter::repeat_n(b'a', crate::proto::MAX_SCHEME_LEN + 1));
1083 assert_eq!(HttpPeekDecision::Reject, state_decision(&oversized_scheme));
1084 }
1085
1086 #[test]
1087 fn test_http1_state_rejects_configured_methods_immediately() {
1088 assert_eq!(HttpPeekDecision::Continue, state_decision(b"PING"));
1089
1090 assert_eq!(
1091 KNOWN_NON_HTTP_PROTOCOL_METHODS,
1092 HttpPeekConfig::new()
1093 .with_known_non_http_protocol_methods()
1094 .skipped_http1_methods,
1095 );
1096
1097 for method in KNOWN_NON_HTTP_PROTOCOL_METHODS {
1098 assert_eq!(
1099 HttpPeekDecision::Reject,
1100 state_decision_with_skipped_methods(
1101 method.as_bytes(),
1102 KNOWN_NON_HTTP_PROTOCOL_METHODS,
1103 ),
1104 "known non-HTTP method {method:?} was not rejected",
1105 );
1106 }
1107 assert_eq!(
1108 HttpPeekDecision::Reject,
1109 state_decision_with_skipped_methods(
1110 b"PINGX / HTTP/1.1\r\n",
1111 KNOWN_NON_HTTP_PROTOCOL_METHODS,
1112 )
1113 );
1114 assert_eq!(
1115 HttpPeekDecision::Continue,
1116 state_decision_with_skipped_methods(b"PIN", KNOWN_NON_HTTP_PROTOCOL_METHODS)
1117 );
1118 assert_eq!(
1119 HttpPeekDecision::Matched(HttpPeekVersion::Http1x),
1120 state_decision_with_skipped_methods(
1121 b"GET / HTTP/1.1\r\n",
1122 KNOWN_NON_HTTP_PROTOCOL_METHODS,
1123 )
1124 );
1125 assert_eq!(
1126 HttpPeekDecision::Reject,
1127 state_decision_with_skipped_methods(b"CUSTOM", &["CUSTOM"])
1128 );
1129 assert_eq!(
1130 HttpPeekDecision::Matched(HttpPeekVersion::H2),
1131 state_decision_with_skipped_methods(H2_MAGIC_PREFIX, &["PRI"])
1132 );
1133 assert_eq!(
1134 HttpPeekDecision::Continue,
1135 state_decision_with_skipped_methods(b"PING", &["", "PI NG", "PÉNG"])
1136 );
1137 }
1138
1139 #[test]
1140 fn test_http1_skipped_method_configuration_is_last_call_wins() {
1141 let config = HttpPeekConfig::new()
1142 .with_known_non_http_protocol_methods()
1143 .with_skipped_http1_methods(&["CUSTOM"]);
1144 assert_eq!(&["CUSTOM"], config.skipped_http1_methods);
1145
1146 let config = HttpPeekConfig::new()
1147 .with_skipped_http1_methods(&["CUSTOM"])
1148 .with_known_non_http_protocol_methods();
1149 assert_eq!(
1150 KNOWN_NON_HTTP_PROTOCOL_METHODS,
1151 config.skipped_http1_methods
1152 );
1153
1154 let router = HttpPeekRouter::new(service_fn(async || Ok::<_, Infallible>("http")))
1155 .with_known_non_http_protocol_methods()
1156 .with_skipped_http1_methods(&["CUSTOM"]);
1157 assert_eq!(&["CUSTOM"], router.peek_config.skipped_http1_methods);
1158 }
1159
1160 #[tokio::test]
1161 async fn test_peek_router() {
1162 let http_service = service_fn(async || Ok::<_, Infallible>("http"));
1163 let fallback_service = service_fn(async || Ok::<_, Infallible>("other"));
1164
1165 let peek_http_svc = HttpPeekRouter::new(http_service).with_fallback(fallback_service);
1166
1167 let response = peek_http_svc
1168 .serve(ServiceInput::new(std::io::Cursor::new(b"".to_vec())))
1169 .await
1170 .unwrap();
1171 assert_eq!("other", response);
1172
1173 let response = peek_http_svc
1174 .serve(ServiceInput::new(std::io::Cursor::new(
1175 b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n".to_vec(),
1176 )))
1177 .await
1178 .unwrap();
1179 assert_eq!("http", response);
1180
1181 let response = peek_http_svc
1182 .serve(ServiceInput::new(std::io::Cursor::new(
1183 b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nfoo".to_vec(),
1184 )))
1185 .await
1186 .unwrap();
1187 assert_eq!("http", response);
1188
1189 const HTTP_METHODS: &[&str] = &[
1190 "GET", "POST", "PUT", "DELETE", "HEAD", "OPTIONS", "CONNECT", "TRACE", "PATCH",
1191 ];
1192 for method in HTTP_METHODS {
1193 let target = if *method == "CONNECT" {
1194 "example.com:443"
1195 } else {
1196 "/foobar"
1197 };
1198 let response = peek_http_svc
1199 .serve(ServiceInput::new(std::io::Cursor::new(
1200 format!("{method} {target} HTTP/1.1\r\n").into_bytes(),
1201 )))
1202 .await
1203 .unwrap();
1204 assert_eq!("http", response);
1205 }
1206
1207 let response = peek_http_svc
1208 .serve(ServiceInput::new(std::io::Cursor::new(b"foo".to_vec())))
1209 .await
1210 .unwrap();
1211 assert_eq!("other", response);
1212
1213 let response = peek_http_svc
1214 .serve(ServiceInput::new(std::io::Cursor::new(b"foobar".to_vec())))
1215 .await
1216 .unwrap();
1217 assert_eq!("other", response);
1218 }
1219
1220 #[tokio::test]
1221 async fn test_peek_http1_connect() {
1222 for timeout in [Some(Duration::from_millis(500)), None] {
1223 let reader = StreamReader::new(
1224 stream_fn(async |mut yielder| {
1225 yielder.yield_item(Bytes::from_static(b"CONN")).await;
1226 tokio::time::sleep(Duration::from_millis(50)).await;
1227 yielder.yield_item(Bytes::from_static(b"EC")).await;
1228 tokio::time::sleep(Duration::from_millis(50)).await;
1229 yielder
1230 .yield_item(Bytes::from_static(b"T foobar.com:443 HTTP/1.1\r\n"))
1231 .await;
1232 })
1233 .map(Ok::<_, std::io::Error>),
1234 );
1235 let writer = tokio::io::sink();
1236
1237 let io = Box::pin(tokio::io::join(reader, writer));
1238
1239 let (http_version, _) = peek_http_input(io, timeout).await.unwrap();
1240
1241 assert_eq!(Some(HttpPeekVersion::Http1x), http_version);
1242 }
1243 }
1244
1245 #[tokio::test]
1246 async fn test_peek_http1_complete_request_line_forms() {
1247 const CASES: &[&[u8]] = &[
1248 b"GET /\r\n",
1249 b"GET http://example.com/legacy\r\n",
1250 b"GET urn:legacy\r\n",
1251 b"GET / HTTP/1.1\r\n",
1252 b"GET /legacy HTTP/1.0\r\n",
1253 b"OPTIONS * HTTP/1.1\r\n",
1254 b"GET http://example.com/resource?q=1 HTTP/1.1\r\n",
1255 b"GET urn:opaque HTTP/1.1\r\n",
1256 b"GET http:/single-slash HTTP/1.1\r\n",
1257 b"CONNECT example.com:443 HTTP/1.1\r\n",
1258 b"PROPFIND /collection HTTP/1.1\r\n",
1259 b"QUERY /search HTTP/1.1\r\n",
1260 b"M-SEARCH /discovery HTTP/1.1\r\n",
1261 b"!#$%&'*+-.^_`|~ /extension HTTP/1.1\r\n",
1262 "GET /café HTTP/1.1\r\n".as_bytes(),
1263 b"GET / HTTP/1.1\n",
1264 b"GET /legacy HTTP/1.0\n",
1265 b"GET / HTTP/1.2\r\n",
1266 b"GET / HTTP/1.9\n",
1267 b"GET /\n",
1268 b"\r\nGET / HTTP/1.1\r\n",
1269 b"\n\n\r\nGET / HTTP/1.1\n",
1270 b"\nCONNECT example.com:443 HTTP/1.1\r\n",
1271 ];
1272
1273 for &content in CASES {
1274 let (version, replayed) = peek_bytes(content, HttpPeekConfig::default()).await;
1275 assert_eq!(Some(HttpPeekVersion::Http1x), version, "{content:?}");
1276 assert_eq!(content, replayed, "{content:?}");
1277 }
1278 }
1279
1280 #[tokio::test]
1281 async fn test_peek_http1_rejects_other_text_protocols_and_invalid_lines() {
1282 const CASES: &[&[u8]] = &[
1283 b"POST /\r\n",
1284 b"get /\r\n",
1285 b"GET *\r\n",
1286 b"GET * HTTP/1.1\r\n",
1287 b"POST * HTTP/1.1\r\n",
1288 b"OPTIONS icap://icap.example.net/service ICAP/1.0\r\n",
1289 b"OPTIONS * RTSP/2.0\r\n",
1290 b"OPTIONS sip:service@example.com SIP/2.0\r\n",
1291 b"GET / HTTP/1.1",
1292 b"GET / HTTP/1.a\r\n",
1293 b"GET / HTTP/1.10\r\n",
1294 b"GET / HTTP/2.0\r\n",
1295 b"GET / HTTP/1.1\r\r",
1296 b"\rGET / HTTP/1.1\r\n",
1297 b"\r\nPRI * HTTP/2.0\r\n\r\nSM\r\n\r\n",
1298 b"GET / HTTP/1.1\r\n",
1299 b"GET /path#fragment HTTP/1.1\r\n",
1300 b"GE(T / HTTP/1.1\r\n",
1301 b"GE(/ HTTP/1.1\r\n",
1302 b" / HTTP/1.1\r\n",
1303 b"GET /bad\ttarget HTTP/1.1\r\n",
1304 b"GET http://[bad]/ HTTP/1.1\r\n",
1305 b"CONNECT /not-authority HTTP/1.1\r\n",
1306 ];
1307
1308 for &content in CASES {
1309 let (version, replayed) = peek_bytes(content, HttpPeekConfig::default()).await;
1310 assert_eq!(None, version, "{content:?}");
1311 assert_eq!(content, replayed, "{content:?}");
1312 }
1313 }
1314
1315 #[tokio::test]
1316 async fn test_peek_handles_single_byte_fragmentation() {
1317 const HTTP1: &[u8] = b"PROPFIND /collection HTTP/1.1\r\nbody";
1318 let (version, replayed) = peek_fragmented_bytes(HTTP1).await;
1319 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1320 assert_eq!(HTTP1, replayed);
1321
1322 const HTTP09: &[u8] = b"GET /legacy\r\n";
1323 let (version, replayed) = peek_fragmented_bytes(HTTP09).await;
1324 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1325 assert_eq!(HTTP09, replayed);
1326
1327 const HTTP1_LENIENT: &[u8] = b"\r\n\nGET /lenient HTTP/1.1\nbody";
1328 let (version, replayed) = peek_fragmented_bytes(HTTP1_LENIENT).await;
1329 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1330 assert_eq!(HTTP1_LENIENT, replayed);
1331
1332 const H2: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nframes";
1333 let (version, replayed) = peek_fragmented_bytes(H2).await;
1334 assert_eq!(Some(HttpPeekVersion::H2), version);
1335 assert_eq!(H2, replayed);
1336
1337 const ICAP: &[u8] = b"OPTIONS icap://icap.example.net/service ICAP/1.0\r\n";
1338 let (version, replayed) = peek_fragmented_bytes(ICAP).await;
1339 assert_eq!(None, version);
1340 assert_eq!(ICAP, replayed);
1341 }
1342
1343 #[tokio::test]
1344 async fn test_peek_http1_configurable_buffer_limits() {
1345 const CONTENT: &[u8] = b"GET /configurable HTTP/1.1\r\n";
1346
1347 let exact = HttpPeekConfig {
1348 max_http1_request_line_size: CONTENT.len(),
1349 read_buffer_size: 1,
1350 ..Default::default()
1351 };
1352
1353 let (version, replayed) = peek_bytes(CONTENT, exact).await;
1354 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1355 assert_eq!(CONTENT, replayed);
1356
1357 let zero_read_size = HttpPeekConfig {
1358 read_buffer_size: 0,
1359 ..exact
1360 };
1361
1362 let (version, replayed) = peek_bytes(CONTENT, zero_read_size).await;
1363 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1364 assert_eq!(CONTENT, replayed);
1365
1366 let too_small = HttpPeekConfig {
1367 max_http1_request_line_size: CONTENT.len() - 1,
1368 ..exact
1369 };
1370
1371 let (version, replayed) = peek_bytes(CONTENT, too_small).await;
1372 assert_eq!(None, version);
1373 assert_eq!(CONTENT, replayed);
1374
1375 let disabled = HttpPeekConfig {
1376 max_http1_request_line_size: 0,
1377 ..exact
1378 };
1379
1380 let (version, replayed) = peek_bytes(CONTENT, disabled).await;
1381 assert_eq!(None, version);
1382 assert_eq!(CONTENT, replayed);
1383
1384 let (version, replayed) = peek_bytes(H2_MAGIC_PREFIX, disabled).await;
1385 assert_eq!(Some(HttpPeekVersion::H2), version);
1386 assert_eq!(H2_MAGIC_PREFIX, replayed);
1387 }
1388
1389 #[tokio::test]
1390 async fn test_peek_http1_starts_reasonably_and_grows_reads_on_demand() {
1391 let mut content = b"PROPFIND /".to_vec();
1392 content.extend(std::iter::repeat_n(
1393 b'a',
1394 INITIAL_HTTP_PEEK_READ_BUFFER_SIZE * 3,
1395 ));
1396 content.extend_from_slice(b" HTTP/1.1\r\nbody");
1397 let read_sizes = Arc::new(Mutex::new(Vec::new()));
1398 let reader = RecordingReader {
1399 inner: std::io::Cursor::new(content),
1400 read_sizes: Arc::clone(&read_sizes),
1401 max_read_size: usize::MAX,
1402 };
1403 let input = tokio::io::join(reader, tokio::io::sink());
1404
1405 let (version, _input) = peek_http_input(input, None).await.unwrap();
1406 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1407
1408 let read_sizes = read_sizes.lock();
1409 assert_eq!(
1410 Some(&INITIAL_HTTP_PEEK_READ_BUFFER_SIZE),
1411 read_sizes.first()
1412 );
1413 assert!(read_sizes.len() >= 3, "read sizes: {read_sizes:?}");
1414 for sizes in read_sizes.windows(2) {
1415 assert!(
1416 sizes[1] <= sizes[0].saturating_mul(2),
1417 "read sizes should grow geometrically: {read_sizes:?}"
1418 );
1419 assert!(
1420 sizes[1] <= DEFAULT_HTTP_PEEK_READ_BUFFER_SIZE,
1421 "read size exceeded configured maximum: {read_sizes:?}"
1422 );
1423 }
1424 }
1425
1426 #[tokio::test]
1427 async fn test_peek_http1_does_not_grow_after_short_reads() {
1428 const CONTENT: &[u8] = b"PROPFIND /fragmented HTTP/1.1\r\n";
1429 let read_sizes = Arc::new(Mutex::new(Vec::new()));
1430 let reader = RecordingReader {
1431 inner: std::io::Cursor::new(CONTENT.to_vec()),
1432 read_sizes: Arc::clone(&read_sizes),
1433 max_read_size: 1,
1434 };
1435 let input = tokio::io::join(reader, tokio::io::sink());
1436
1437 let (version, _input) = peek_http_input(input, None).await.unwrap();
1438 assert_eq!(Some(HttpPeekVersion::Http1x), version);
1439
1440 let read_sizes = read_sizes.lock();
1441 assert!(read_sizes.len() > 1);
1442 assert!(
1443 read_sizes
1444 .iter()
1445 .all(|&size| size == INITIAL_HTTP_PEEK_READ_BUFFER_SIZE),
1446 "short reads should not grow the next read window: {read_sizes:?}"
1447 );
1448 }
1449
1450 #[tokio::test]
1451 async fn test_peek_router_skips_configured_idle_method() {
1452 async fn fallback_service_fn(
1453 mut stream: impl Io + Unpin,
1454 ) -> Result<&'static str, io::Error> {
1455 let mut method = [0_u8; 4];
1456 stream.read_exact(&mut method).await?;
1457 assert_eq!(b"PING", &method);
1458 Ok("fallback")
1459 }
1460
1461 let http_service = service_fn(async || Ok::<_, Infallible>("http"));
1462 let router = HttpPeekRouter::new(http_service)
1463 .with_known_non_http_protocol_methods()
1464 .with_fallback(service_fn(fallback_service_fn));
1465 let input = tokio::io::join(IdleAfterPrefix::new(b"PING"), tokio::io::sink());
1466
1467 let result = tokio::time::timeout(Duration::from_secs(1), router.serve(input))
1468 .await
1469 .expect("configured method must fall back without waiting for a delimiter")
1470 .unwrap();
1471 assert_eq!("fallback", result);
1472 }
1473
1474 #[tokio::test]
1475 async fn test_peek_http1_timeout_replays_partial_request_line() {
1476 const PREFIX: &[u8] = b"GET /slow ";
1477 const SUFFIX: &[u8] = b"HTTP/1.1\r\n";
1478 let reader = StreamReader::new(
1479 stream_fn(async |mut yielder| {
1480 yielder.yield_item(Bytes::from_static(PREFIX)).await;
1481 tokio::time::sleep(Duration::from_millis(50)).await;
1482 yielder.yield_item(Bytes::from_static(SUFFIX)).await;
1483 })
1484 .map(Ok::<_, std::io::Error>),
1485 );
1486 let io = Box::pin(tokio::io::join(reader, tokio::io::sink()));
1487
1488 let (version, mut io) = peek_http_input(io, Some(Duration::from_millis(10)))
1489 .await
1490 .unwrap();
1491 assert_eq!(None, version);
1492
1493 let mut replayed = Vec::new();
1494 io.read_to_end(&mut replayed).await.unwrap();
1495 assert_eq!([PREFIX, SUFFIX].concat(), replayed);
1496 }
1497
1498 #[tokio::test]
1499 async fn test_http_router_default_fail_open_replays_partial_request_line_to_fallback() {
1500 const PREFIX: &[u8] = b"GET /slow ";
1501
1502 async fn fallback_service_fn(
1503 mut stream: impl Io + Unpin,
1504 ) -> Result<&'static str, io::Error> {
1505 let mut prefix = [0_u8; PREFIX.len()];
1506 stream.read_exact(&mut prefix).await?;
1507 assert_eq!(PREFIX, prefix);
1508 Ok("fallback")
1509 }
1510
1511 let router = HttpPeekRouter::new(service_fn(async || Ok::<_, Infallible>("http")))
1512 .with_peek_timeout(Duration::from_millis(10))
1513 .with_fallback(service_fn(fallback_service_fn));
1514 let input = tokio::io::join(IdleAfterPrefix::new(PREFIX), tokio::io::sink());
1515
1516 let result = router.serve(input).await.unwrap();
1517 assert_eq!("fallback", result);
1518 assert_eq!(PeekTimeoutPolicy::FailOpen, PeekTimeoutPolicy::default());
1519 }
1520
1521 #[tokio::test]
1522 async fn test_all_http_router_modes_fail_closed_without_invoking_fallback() {
1523 let fallback_calls = Arc::new(AtomicUsize::new(0));
1524
1525 macro_rules! assert_fail_closed {
1526 ($router:expr) => {{
1527 let fallback_calls_for_service = Arc::clone(&fallback_calls);
1528 let fallback = service_fn(move || {
1529 fallback_calls_for_service.fetch_add(1, Ordering::SeqCst);
1530 async { Ok::<_, Infallible>("fallback") }
1531 });
1532 let router = $router
1533 .with_peek_timeout(Duration::from_millis(10))
1534 .with_peek_timeout_policy(PeekTimeoutPolicy::FailClosed)
1535 .with_fallback(fallback);
1536 let input =
1537 tokio::io::join(IdleAfterPrefix::new(b"GET /fragmented "), tokio::io::sink());
1538
1539 let error = router.serve(input).await.unwrap_err();
1540 assert!(error.downcast_ref::<PeekTimeoutError>().is_some());
1541 assert_eq!(0, fallback_calls.load(Ordering::SeqCst));
1542 }};
1543 }
1544
1545 assert_fail_closed!(HttpPeekRouter::new(service_fn(async || {
1546 Ok::<_, Infallible>("http")
1547 })));
1548 assert_fail_closed!(HttpPeekRouter::new_http1(service_fn(async || {
1549 Ok::<_, Infallible>("http1")
1550 })));
1551 assert_fail_closed!(HttpPeekRouter::new_h2(service_fn(async || {
1552 Ok::<_, Infallible>("h2")
1553 })));
1554 assert_fail_closed!(HttpPeekRouter::new_dual(
1555 service_fn(async || Ok::<_, Infallible>("http1")),
1556 service_fn(async || Ok::<_, Infallible>("h2")),
1557 ));
1558 }
1559
1560 #[tokio::test]
1561 async fn test_http_definitive_mismatch_falls_back_under_both_timeout_policies() {
1562 for policy in [PeekTimeoutPolicy::FailOpen, PeekTimeoutPolicy::FailClosed] {
1563 let router = HttpPeekRouter::new(service_fn(async || Ok::<_, Infallible>("http")))
1564 .with_peek_timeout(Duration::from_millis(10))
1565 .with_peek_timeout_policy(policy)
1566 .with_fallback(service_fn(async || Ok::<_, Infallible>("fallback")));
1567
1568 let response = router
1569 .serve(ServiceInput::new(std::io::Cursor::new(vec![0])))
1570 .await
1571 .unwrap();
1572 assert_eq!("fallback", response);
1573 }
1574 }
1575
1576 #[tokio::test]
1577 async fn test_peek_http1_router() {
1578 let http_service = service_fn(async || Ok::<_, Infallible>("http1"));
1579 let fallback_service = service_fn(async || Ok::<_, Infallible>("other"));
1580
1581 let peek_http_svc = HttpPeekRouter::new_http1(http_service).with_fallback(fallback_service);
1582
1583 let response = peek_http_svc
1584 .serve(ServiceInput::new(std::io::Cursor::new(b"".to_vec())))
1585 .await
1586 .unwrap();
1587 assert_eq!("other", response);
1588
1589 let response = peek_http_svc
1590 .serve(ServiceInput::new(std::io::Cursor::new(
1591 b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nfoo".to_vec(),
1592 )))
1593 .await
1594 .unwrap();
1595 assert_eq!("other", response);
1596
1597 const HTTP_METHODS: &[&str] = &[
1598 "GET", "POST", "PUT", "DELETE", "HEAD", "OPTIONS", "CONNECT", "TRACE", "PATCH",
1599 ];
1600 for method in HTTP_METHODS {
1601 let target = if *method == "CONNECT" {
1602 "example.com:443"
1603 } else {
1604 "/foobar"
1605 };
1606 let response = peek_http_svc
1607 .serve(ServiceInput::new(std::io::Cursor::new(
1608 format!("{method} {target} HTTP/1.1\r\n").into_bytes(),
1609 )))
1610 .await
1611 .unwrap();
1612 assert_eq!("http1", response);
1613 }
1614
1615 let response = peek_http_svc
1616 .serve(ServiceInput::new(std::io::Cursor::new(b"foo".to_vec())))
1617 .await
1618 .unwrap();
1619 assert_eq!("other", response);
1620
1621 let response = peek_http_svc
1622 .serve(ServiceInput::new(std::io::Cursor::new(b"foobar".to_vec())))
1623 .await
1624 .unwrap();
1625 assert_eq!("other", response);
1626 }
1627
1628 #[tokio::test]
1629 async fn test_peek_h2_router() {
1630 let http_service = service_fn(async || Ok::<_, Infallible>("h2"));
1631 let fallback_service = service_fn(async || Ok::<_, Infallible>("other"));
1632
1633 let peek_http_svc = HttpPeekRouter::new_h2(http_service).with_fallback(fallback_service);
1634
1635 let response = peek_http_svc
1636 .serve(ServiceInput::new(std::io::Cursor::new(b"".to_vec())))
1637 .await
1638 .unwrap();
1639 assert_eq!("other", response);
1640
1641 let response = peek_http_svc
1642 .serve(ServiceInput::new(std::io::Cursor::new(
1643 b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nfoo".to_vec(),
1644 )))
1645 .await
1646 .unwrap();
1647 assert_eq!("h2", response);
1648
1649 const HTTP_METHODS: &[&str] = &[
1650 "GET", "POST", "PUT", "DELETE", "HEAD", "OPTIONS", "CONNECT", "TRACE", "PATCH",
1651 ];
1652 for method in HTTP_METHODS {
1653 let target = if *method == "CONNECT" {
1654 "example.com:443"
1655 } else {
1656 "/foobar"
1657 };
1658 let response = peek_http_svc
1659 .serve(ServiceInput::new(std::io::Cursor::new(
1660 format!("{method} {target} HTTP/1.1\r\n").into_bytes(),
1661 )))
1662 .await
1663 .unwrap();
1664 assert_eq!("other", response);
1665 }
1666
1667 let response = peek_http_svc
1668 .serve(ServiceInput::new(std::io::Cursor::new(b"foo".to_vec())))
1669 .await
1670 .unwrap();
1671 assert_eq!("other", response);
1672
1673 let response = peek_http_svc
1674 .serve(ServiceInput::new(std::io::Cursor::new(b"foobar".to_vec())))
1675 .await
1676 .unwrap();
1677 assert_eq!("other", response);
1678 }
1679
1680 #[tokio::test]
1681 async fn test_peek_dual_router() {
1682 let http1_service = service_fn(async || Ok::<_, Infallible>("http1"));
1683 let h2_service = service_fn(async || Ok::<_, Infallible>("h2"));
1684 let fallback_service = service_fn(async || Ok::<_, Infallible>("other"));
1685
1686 let peek_http_svc =
1687 HttpPeekRouter::new_dual(http1_service, h2_service).with_fallback(fallback_service);
1688
1689 let response = peek_http_svc
1690 .serve(ServiceInput::new(std::io::Cursor::new(b"".to_vec())))
1691 .await
1692 .unwrap();
1693 assert_eq!("other", response);
1694
1695 let response = peek_http_svc
1696 .serve(ServiceInput::new(std::io::Cursor::new(
1697 b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nfoo".to_vec(),
1698 )))
1699 .await
1700 .unwrap();
1701 assert_eq!("h2", response);
1702
1703 const HTTP_METHODS: &[&str] = &[
1704 "GET", "POST", "PUT", "DELETE", "HEAD", "OPTIONS", "CONNECT", "TRACE", "PATCH",
1705 ];
1706 for method in HTTP_METHODS {
1707 let target = if *method == "CONNECT" {
1708 "example.com:443"
1709 } else {
1710 "/foobar"
1711 };
1712 let response = peek_http_svc
1713 .serve(ServiceInput::new(std::io::Cursor::new(
1714 format!("{method} {target} HTTP/1.1\r\n").into_bytes(),
1715 )))
1716 .await
1717 .unwrap();
1718 assert_eq!("http1", response);
1719 }
1720
1721 let response = peek_http_svc
1722 .serve(ServiceInput::new(std::io::Cursor::new(b"foo".to_vec())))
1723 .await
1724 .unwrap();
1725 assert_eq!("other", response);
1726
1727 let response = peek_http_svc
1728 .serve(ServiceInput::new(std::io::Cursor::new(b"foobar".to_vec())))
1729 .await
1730 .unwrap();
1731 assert_eq!("other", response);
1732 }
1733
1734 #[tokio::test]
1735 async fn test_peek_router_read_eof() {
1736 const CONTENT: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\nfoobar";
1737
1738 async fn http_service_fn(mut stream: impl Io + Unpin) -> Result<&'static str, BoxError> {
1739 let mut v = Vec::default();
1740 _ = stream.read_to_end(&mut v).await?;
1741 assert_eq!(CONTENT, v);
1742
1743 Ok("ok")
1744 }
1745 let http_service = service_fn(http_service_fn);
1746
1747 let peek_http_svc = HttpPeekRouter::new(http_service).with_fallback(RejectService::<
1748 &'static str,
1749 RejectError,
1750 >::new(
1751 RejectError::default(),
1752 ));
1753
1754 let response = peek_http_svc
1755 .serve(ServiceInput::new(std::io::Cursor::new(CONTENT.to_vec())))
1756 .await
1757 .unwrap();
1758 assert_eq!("ok", response);
1759 }
1760
1761 #[tokio::test]
1762 async fn test_peek_router_read_no_http_eof() {
1763 let cases = [
1764 "",
1765 "foo",
1766 "abcd",
1767 "abcde",
1768 "foobarbazbananas",
1769 "Lorem ipsum dolor sit amet, consectetur adipiscing elit. Nunc vehicula turpis nibh, eget euismod enim elementum et.",
1770 ];
1771 for content in cases {
1772 async fn http_service_fn() -> Result<Vec<u8>, BoxError> {
1773 Ok("http".as_bytes().to_vec())
1774 }
1775 let http_service = service_fn(http_service_fn);
1776
1777 async fn other_service_fn(mut stream: impl Io + Unpin) -> Result<Vec<u8>, BoxError> {
1778 let mut v = Vec::default();
1779 _ = stream.read_to_end(&mut v).await?;
1780 Ok(v)
1781 }
1782 let other_service = service_fn(other_service_fn);
1783
1784 let peek_http_svc = HttpPeekRouter::new(http_service).with_fallback(other_service);
1785
1786 let response = peek_http_svc
1787 .serve(ServiceInput::new(std::io::Cursor::new(
1788 content.as_bytes().to_vec(),
1789 )))
1790 .await
1791 .unwrap();
1792
1793 assert_eq!(content.as_bytes(), &response[..]);
1794 }
1795 }
1796
1797 struct RecordingReader {
1798 inner: std::io::Cursor<Vec<u8>>,
1799 read_sizes: Arc<Mutex<Vec<usize>>>,
1800 max_read_size: usize,
1801 }
1802
1803 impl AsyncRead for RecordingReader {
1804 fn poll_read(
1805 mut self: Pin<&mut Self>,
1806 _cx: &mut Context<'_>,
1807 buffer: &mut ReadBuf<'_>,
1808 ) -> Poll<io::Result<()>> {
1809 self.read_sizes.lock().push(buffer.remaining());
1810 let start = self.inner.position() as usize;
1811 let end = (start + buffer.remaining().min(self.max_read_size))
1812 .min(self.inner.get_ref().len());
1813 if start < end {
1814 buffer.put_slice(&self.inner.get_ref()[start..end]);
1815 self.inner.set_position(end as u64);
1816 }
1817 Poll::Ready(Ok(()))
1818 }
1819 }
1820
1821 struct IdleAfterPrefix {
1822 prefix: Option<&'static [u8]>,
1823 }
1824
1825 impl IdleAfterPrefix {
1826 fn new(prefix: &'static [u8]) -> Self {
1827 Self {
1828 prefix: Some(prefix),
1829 }
1830 }
1831 }
1832
1833 impl AsyncRead for IdleAfterPrefix {
1834 fn poll_read(
1835 mut self: Pin<&mut Self>,
1836 _cx: &mut Context<'_>,
1837 buffer: &mut ReadBuf<'_>,
1838 ) -> Poll<io::Result<()>> {
1839 if let Some(prefix) = self.prefix.take() {
1840 buffer.put_slice(prefix);
1841 Poll::Ready(Ok(()))
1842 } else {
1843 Poll::Pending
1844 }
1845 }
1846 }
1847}