rama_net/client/
connect_request.rs1use 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)]
17pub 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 #[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 #[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 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 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}