Skip to main content

rama_net/http/server/
peek.rs

1//! types and logic for [`HttpPeekRouter`]
2
3use 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
24/// Default maximum number of bytes inspected for an HTTP/1 request-line.
25///
26/// RFC 9112 recommends that HTTP senders and recipients support request-lines
27/// of at least 8000 octets. The slightly larger power-of-two default leaves
28/// room for the terminating CRLF while keeping protocol detection bounded.
29pub const DEFAULT_HTTP1_REQUEST_LINE_MAX_SIZE: usize = kib(8);
30
31/// Default maximum number of bytes requested from the transport per peek read.
32pub const DEFAULT_HTTP_PEEK_READ_BUFFER_SIZE: usize = 512;
33
34/// Known non-HTTP protocol prefixes which resemble HTTP/1 methods.
35///
36/// Covers Redis, IRC, SMTP, FTP, SSH and HAProxy PROXY v1. Matching is
37/// case-sensitive and fail-fast, so it can also reject HTTP extension methods
38/// with one of these prefixes. Server-first protocols are only recognized once
39/// a listed client prefix arrives.
40pub const KNOWN_NON_HTTP_PROTOCOL_METHODS: &[&str] =
41    &["PING", "EHLO", "HELO", "USER", "NICK", "SSH", "PROXY"];
42
43/// Initial HTTP peek read size.
44///
45/// This covers the overwhelming majority of HTTP/1 request lines in one
46/// transport read while retaining geometric growth for unusually long targets.
47const INITIAL_HTTP_PEEK_READ_BUFFER_SIZE: usize = 128;
48
49/// Configuration used while detecting HTTP on a byte stream.
50#[non_exhaustive]
51#[derive(Debug, Clone, Copy, PartialEq, Eq)]
52pub struct HttpPeekConfig {
53    /// timeout applied to the complete HTTP peek operation
54    pub timeout: Option<Duration>,
55    /// policy applied when the timeout expires before HTTP detection reaches a
56    /// definitive match or rejection
57    pub timeout_policy: PeekTimeoutPolicy,
58    /// maximum number of bytes inspected for an HTTP/1 request-line
59    pub max_http1_request_line_size: usize,
60    /// maximum number of bytes requested per transport read
61    pub read_buffer_size: usize,
62    /// HTTP/1 method prefixes routed to the fallback before a delimiter.
63    ///
64    /// Matching is case-sensitive and fail-fast, so `PING` also skips `PINGX`.
65    /// Empty or invalid method tokens cannot match.
66    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    /// Create an HTTP peek configuration using the defaults.
83    #[must_use]
84    pub fn new() -> Self {
85        Self::default()
86    }
87
88    rama_utils::macros::generate_set_and_with! {
89        /// Set how an inconclusive HTTP peek timeout is handled.
90        ///
91        /// Defaults to [`PeekTimeoutPolicy::FailOpen`]. A definitive non-HTTP
92        /// prefix remains a normal non-match under either policy.
93        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        /// Use [`KNOWN_NON_HTTP_PROTOCOL_METHODS`] as the skipped method list.
101        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        /// Set [`HttpPeekConfig::skipped_http1_methods`].
109        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/// A [`Service`] router that can be used to support
120/// http/1x and h2 traffic as well as non-tls traffic.
121///
122/// By default non-http traffic is rejected using [`RejectService`].
123/// Use [`HttpPeekRouter::with_fallback`] to configure the fallback service.
124/// A definitive non-HTTP prefix invokes that fallback under either timeout
125/// policy; an inconclusive timeout is fail-open by default and can be made
126/// fail-closed with [`HttpPeekRouter::with_peek_timeout_policy`].
127#[derive(Debug, Clone)]
128pub struct HttpPeekRouter<T, F = RejectService<(), NoHttpRejectError>> {
129    http_acceptor: T,
130    fallback: F,
131    peek_config: HttpPeekConfig,
132}
133
134/// Type wrapper used by [`HttpPeekRouter::new_dual`]
135/// to serve http/1x and h2 separately.
136#[derive(Debug, Clone)]
137pub struct HttpDualAcceptor<T, U> {
138    http1: T,
139    h2: U,
140}
141
142/// Type wrapper used by [`HttpPeekRouter::new`]
143/// to serve http/1x and h2 with a single service.
144#[derive(Debug, Clone)]
145pub struct HttpAutoAcceptor<T>(T);
146
147/// Type wrapper used by [`HttpPeekRouter::new_http1`]
148/// to only serve http/1x, and send h2 to the fallback.
149#[derive(Debug, Clone)]
150pub struct Http1Acceptor<T>(T);
151
152/// Type wrapper used by [`HttpPeekRouter::new_h2`]
153/// to only serve h2, and send http/1x to the fallback.
154#[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    /// Create a new [`HttpPeekRouter`] using a service
164    /// which can handle h2 and http/1x versions alike.
165    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    /// Create a new [`HttpPeekRouter`] using a service
176    /// which handles http/1x traffic but forwards h2 traffic to fallback.
177    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    /// Create a new [`HttpPeekRouter`] using a service
188    /// which handles h2 traffic but forwards http/1x traffic to fallback.
189    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    /// Attach a fallback [`Service`].
200    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        /// Use [`KNOWN_NON_HTTP_PROTOCOL_METHODS`] as the skipped method list.
212        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        /// Set the maximum time spent peeking for HTTP.
220        ///
221        /// A timeout is inconclusive. Use
222        /// [`HttpPeekRouter::with_peek_timeout_policy`] to choose whether it
223        /// invokes the fallback or rejects the connection.
224        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        /// Set how an inconclusive HTTP peek timeout is handled.
232        ///
233        /// Defaults to [`PeekTimeoutPolicy::FailOpen`]. A definitive non-HTTP
234        /// prefix invokes the fallback under either policy.
235        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        /// Set the configuration used while peeking for HTTP.
243        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        /// Set [`HttpPeekConfig::skipped_http1_methods`].
251        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    /// Create a new [`HttpPeekRouter`] using a service
263    /// which handles http/1x and h2 in two separate services.
264    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    /// A leading `\r` was read; only the `\n` completing an empty line may follow.
440    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        // `HTTP/1.` + any minor digit + line terminator.
578        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                    // httparse parity: skip empty lines before the request-line
611                    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                                    // httparse parity: bare LF terminates the simple-request
661                                    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                    // httparse parity: bare LF may terminate the request-line
746                    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    // Method bytes are RFC 9110 `tchar`, and target validation guarantees
775    // UTF-8, so both views are safe and borrow the single replay buffer.
776    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
802/// Detect HTTP using [`HttpPeekConfig::default`] and return an input that
803/// replays every byte consumed during detection.
804///
805/// Timeouts use [`PeekTimeoutPolicy::FailOpen`] for backward compatibility. Use
806/// [`peek_http_input_with_timeout_policy`] to select an explicit policy.
807pub 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
823/// Detect HTTP with an explicit timeout policy and return an input that replays
824/// every byte consumed during detection.
825///
826/// A definitive non-HTTP prefix returns `Ok((None, input))` under either policy.
827/// An inconclusive timeout returns [`PeekTimeoutError`] when
828/// [`PeekTimeoutPolicy::FailClosed`] is selected.
829pub 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
854/// Detect HTTP using explicit resource limits and return an input that
855/// replays every byte consumed during detection.
856///
857/// HTTP/1.0 and HTTP/1.1 are accepted only after a complete request-line
858/// containing a valid method token, request target, `HTTP/1.<digit>` version, and
859/// line terminator has been read. An HTTP/0.9 simple-request is accepted as
860/// HTTP/1x after `GET`, a valid request target, and a line terminator.
861/// Matching lenient HTTP/1 parsers such as httparse, a bare LF is accepted
862/// as line terminator and empty lines before the request-line are skipped
863/// (they count toward [`HttpPeekConfig::max_http1_request_line_size`]).
864/// HTTP/2 is accepted only after its complete client preface, starting at
865/// the first byte.
866pub 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
968/// [`PrefixedIo`] alias used by [`HttpPeekRouter`].
969pub 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        // A leading byte of a potentially valid UTF-8 target is not enough
1069        // information to reject; validation completes at the target delimiter.
1070        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}