Skip to main content

rama_net/client/
connect_request.rs

1use crate::{
2    AuthorityInputExt, Protocol, ProtocolInputExt, TransportProtocolInputExt,
3    address::{HostWithOptPort, HostWithPort},
4    transport::TransportProtocol,
5};
6
7use rama_core::{Fork, extensions::Extensions, extensions::ExtensionsRef};
8
9#[cfg(feature = "http")]
10use crate::{
11    HttpVersionInputExt, TargetHttpVersionInputExt,
12    http::{HttpRequestVersion, TargetHttpVersion, Version},
13};
14
15#[non_exhaustive]
16#[derive(Debug, Clone)]
17/// A protocol-independent request to establish a client connection.
18pub struct ConnectRequest {
19    pub authority: HostWithPort,
20    pub extensions: Extensions,
21    pub application_protocol: Option<Protocol>,
22    pub transport_protocol: Option<TransportProtocol>,
23}
24
25impl ConnectRequest {
26    /// Create a new [`ConnectRequest`] with default [`Extensions`].
27    #[must_use]
28    pub fn new(authority: HostWithPort) -> Self {
29        Self {
30            authority,
31            extensions: Extensions::new(),
32            application_protocol: None,
33            transport_protocol: None,
34        }
35    }
36
37    /// Create a new [`ConnectRequest`] with given [`Extensions`].
38    #[must_use]
39    pub const fn new_with_extensions(authority: HostWithPort, extensions: Extensions) -> Self {
40        Self {
41            authority,
42            extensions,
43            application_protocol: None,
44            transport_protocol: None,
45        }
46    }
47
48    rama_utils::macros::generate_set_and_with! {
49        /// Define the application [`Protocol`] to this [`ConnectRequest`]
50        /// requested for this connection.
51        ///
52        /// By default the flow context will define the used application protocol.
53        pub fn application_protocol(mut self, protocol: Option<Protocol>) -> Self {
54            self.application_protocol = protocol;
55            self
56        }
57    }
58
59    rama_utils::macros::generate_set_and_with! {
60        /// Define the [`TransportProtocol`] to this [`ConnectRequest`]
61        /// requested for this connection.
62        ///
63        /// By default it will defined by the flow receiver itself.
64        pub fn transport_protocol(mut self, protocol: Option<TransportProtocol>) -> Self {
65            self.transport_protocol = protocol;
66            self
67        }
68    }
69}
70
71impl Fork for ConnectRequest {
72    fn fork(&self) -> Self {
73        Self {
74            authority: self.authority.clone(),
75            extensions: self.extensions.fork(),
76            application_protocol: self.application_protocol.clone(),
77            transport_protocol: self.transport_protocol,
78        }
79    }
80}
81
82impl ExtensionsRef for ConnectRequest {
83    fn extensions(&self) -> &Extensions {
84        &self.extensions
85    }
86}
87
88impl AuthorityInputExt for ConnectRequest {
89    fn authority(&self) -> Option<HostWithOptPort> {
90        Some(self.authority.clone().into())
91    }
92}
93
94impl ProtocolInputExt for ConnectRequest {
95    fn protocol(&self) -> Option<&Protocol> {
96        self.application_protocol.as_ref()
97    }
98}
99
100impl TransportProtocolInputExt for ConnectRequest {
101    fn transport_protocol(&self) -> Option<TransportProtocol> {
102        self.transport_protocol
103    }
104}
105
106#[cfg(feature = "http")]
107impl TargetHttpVersionInputExt for ConnectRequest {
108    fn target_http_version(&self) -> Option<Version> {
109        self.extensions
110            .get_ref::<TargetHttpVersion>()
111            .map(|target| target.0)
112    }
113}
114
115#[cfg(feature = "http")]
116impl HttpVersionInputExt for ConnectRequest {
117    fn http_version(&self) -> Option<Version> {
118        self.extensions
119            .get_ref::<HttpRequestVersion>()
120            .map(|version| version.0)
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use rama_core::extensions::Extension;
127
128    use super::*;
129
130    #[derive(Debug, Extension)]
131    struct AttemptMarker;
132
133    #[test]
134    fn fork_isolates_attempt_extensions() {
135        let request = ConnectRequest::new(HostWithPort::example_domain_https());
136        let attempt = request.fork();
137
138        attempt.extensions.insert(AttemptMarker);
139
140        assert!(!request.extensions.contains::<AttemptMarker>());
141        assert!(attempt.extensions.contains::<AttemptMarker>());
142    }
143
144    #[cfg(feature = "http")]
145    #[test]
146    fn target_http_version_comes_from_extension() {
147        let request = ConnectRequest::new(HostWithPort::example_domain_https());
148        assert_eq!(request.target_http_version(), None);
149
150        request
151            .extensions
152            .insert(TargetHttpVersion(Version::HTTP_2));
153
154        assert_eq!(request.target_http_version(), Some(Version::HTTP_2));
155    }
156
157    #[cfg(feature = "http")]
158    #[test]
159    fn request_http_version_is_not_a_target_requirement() {
160        let request = ConnectRequest::new(HostWithPort::example_domain_https());
161        request
162            .extensions
163            .insert(HttpRequestVersion(Version::HTTP_2));
164
165        assert_eq!(request.http_version(), Some(Version::HTTP_2));
166        assert_eq!(request.target_http_version(), None);
167    }
168}