rama_http/layer/version_adapter/
response.rs1use rama_core::error::{BoxError, ErrorContext as _};
2use rama_core::telemetry::tracing;
3use rama_core::{Layer, Service};
4use rama_http_headers::{Connection, HeaderMapExt, SecWebSocketAccept, SecWebSocketKey, Upgrade};
5use rama_http_types::header::SEC_WEBSOCKET_ACCEPT;
6use rama_http_types::proto::h2::ext::Protocol;
7use rama_http_types::{Request, Response, StatusCode, Version};
8
9use super::request::{is_websocket_protocol, request_connect_protocol};
10use crate::layer::remove_header::remove_illegal_h2_response_headers;
11
12#[derive(Clone, Debug)]
13pub struct ResponseVersionAdapter<S> {
23 inner: S,
24}
25
26impl<S> ResponseVersionAdapter<S> {
27 pub fn new(inner: S) -> Self {
28 Self { inner }
29 }
30}
31
32impl<S, Body> Service<Request<Body>> for ResponseVersionAdapter<S>
33where
34 S: Service<Request<Body>, Output = Response, Error: Into<BoxError>>,
35 Body: Send + 'static,
36{
37 type Output = S::Output;
38 type Error = BoxError;
39
40 async fn serve(&self, req: Request<Body>) -> Result<Self::Output, Self::Error> {
41 let request_ctx = ResponseVersionAdaptCtx::from_request(&req);
42
43 let mut resp = self.inner.serve(req).await.into_box_error()?;
44 adapt_response_version(&mut resp, &request_ctx)?;
45
46 Ok(resp)
47 }
48}
49
50#[non_exhaustive]
51#[derive(Clone, Debug, Default)]
52pub struct ResponseVersionAdapterLayer;
57
58impl<S> Layer<S> for ResponseVersionAdapterLayer {
59 type Service = ResponseVersionAdapter<S>;
60
61 fn layer(&self, inner: S) -> Self::Service {
62 ResponseVersionAdapter { inner }
63 }
64}
65
66#[derive(Debug, Clone, Default)]
74pub struct ResponseVersionAdaptCtx {
75 pub version: Version,
77 pub connect_protocol: Option<Protocol>,
81 pub websocket_key: Option<SecWebSocketKey>,
84}
85
86impl ResponseVersionAdaptCtx {
87 pub fn from_request<Body>(request: &Request<Body>) -> Self {
89 Self {
90 version: request.version(),
91 connect_protocol: request_connect_protocol(request),
92 websocket_key: request.headers().typed_get::<SecWebSocketKey>(),
93 }
94 }
95
96 fn is_websocket(&self) -> bool {
97 self.connect_protocol
98 .as_ref()
99 .is_some_and(is_websocket_protocol)
100 }
101}
102
103pub fn adapt_response_version<Body>(
106 response: &mut Response<Body>,
107 request_ctx: &ResponseVersionAdaptCtx,
108) -> Result<(), BoxError> {
109 let resp_version = response.version();
110 if resp_version == request_ctx.version {
111 tracing::trace!(
112 version = ?response.version(),
113 "response version is already correct, no version switching needed",
114 );
115 return Ok(());
116 }
117
118 tracing::trace!(
119 ?resp_version,
120 target_version = ?request_ctx.version,
121 "changing response version",
122 );
123
124 let resp_is_h1 = resp_version <= Version::HTTP_11;
128 let target_is_h1 = request_ctx.version <= Version::HTTP_11;
129
130 match (resp_is_h1, target_is_h1) {
131 (true, false) => upgrade_response_to_h2_or_h3(response, request_ctx)?,
132 (false, true) => downgrade_response_to_h1(response, request_ctx)?,
133 (true, true) | (false, false) => {}
135 }
136
137 *response.version_mut() = request_ctx.version;
138 Ok(())
139}
140
141fn upgrade_response_to_h2_or_h3<Body>(
154 response: &mut Response<Body>,
155 request_ctx: &ResponseVersionAdaptCtx,
156) -> Result<(), BoxError> {
157 if response.status() == StatusCode::SWITCHING_PROTOCOLS {
158 if request_ctx.is_websocket() {
159 tracing::trace!("translating h1 websocket 101 response into h2/h3 200 OK");
160 *response.status_mut() = StatusCode::OK;
161 response.headers_mut().remove(SEC_WEBSOCKET_ACCEPT);
164 } else {
165 return Err(BoxError::from(format!(
166 "cannot translate a `101 Switching Protocols` response to HTTP/2+ for protocol {}: only websocket is supported",
167 request_ctx
168 .connect_protocol
169 .as_ref()
170 .map_or("<unknown upgrade>", Protocol::as_str),
171 )));
172 }
173 }
174
175 remove_illegal_h2_response_headers(response.headers_mut());
179 Ok(())
180}
181
182fn downgrade_response_to_h1<Body>(
190 response: &mut Response<Body>,
191 request_ctx: &ResponseVersionAdaptCtx,
192) -> Result<(), BoxError> {
193 let Some(protocol) = request_ctx.connect_protocol.as_ref() else {
194 return Ok(());
196 };
197 if !is_websocket_protocol(protocol) {
198 return Err(BoxError::from(format!(
199 "cannot translate an Extended CONNECT `{}` response to HTTP/1: only websocket is supported",
200 protocol.as_str(),
201 )));
202 }
203
204 if response.status() == StatusCode::OK {
207 tracing::trace!("translating h2/h3 websocket 200 response into h1 101 Switching Protocols");
208 *response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
209
210 let headers = response.headers_mut();
211 headers.typed_insert(Upgrade::websocket());
212 headers.typed_insert(Connection::upgrade());
213 if let Some(key) = request_ctx.websocket_key.clone() {
214 let accept = SecWebSocketAccept::try_from(key)
215 .context("derive Sec-WebSocket-Accept for h1 websocket handshake response")?;
216 headers.typed_insert(accept);
217 } else {
218 tracing::debug!(
219 "no Sec-WebSocket-Key captured; emitting h1 websocket 101 without Sec-WebSocket-Accept",
220 );
221 }
222 }
223 Ok(())
224}
225
226#[cfg(test)]
227mod tests {
228 use super::*;
229 use rama_core::extensions::ExtensionsRef;
230 use rama_http_types::Method;
231 use rama_http_types::header::{CONNECTION, SEC_WEBSOCKET_KEY, TRANSFER_ENCODING, UPGRADE};
232
233 const SAMPLE_KEY: &str = "dGhlIHNhbXBsZSBub25jZQ==";
234 const SAMPLE_ACCEPT: &str = "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=";
236
237 fn ctx_with_version(version: Version) -> ResponseVersionAdaptCtx {
238 ResponseVersionAdaptCtx {
239 version,
240 ..Default::default()
241 }
242 }
243
244 fn websocket_ctx(version: Version) -> ResponseVersionAdaptCtx {
245 let req = Request::builder()
246 .version(version)
247 .uri("https://example.com/chat")
248 .header(UPGRADE, "websocket")
249 .header(CONNECTION, "Upgrade")
250 .header(SEC_WEBSOCKET_KEY, SAMPLE_KEY)
251 .body(())
252 .unwrap();
253 ResponseVersionAdaptCtx::from_request(&req)
254 }
255
256 fn connect_udp_ctx(version: Version) -> ResponseVersionAdaptCtx {
257 let req = Request::builder()
258 .version(version)
259 .method(Method::CONNECT)
260 .uri("https://example.com/.well-known/masque/udp/1.2.3.4/443/")
261 .body(())
262 .unwrap();
263 req.extensions()
264 .insert(Protocol::from_static("connect-udp"));
265 ResponseVersionAdaptCtx::from_request(&req)
266 }
267
268 #[test]
269 fn test_h1_to_h2_strips_hop_by_hop_headers() {
270 let mut resp = Response::builder()
271 .version(Version::HTTP_11)
272 .header(CONNECTION, "keep-alive")
273 .header("keep-alive", "timeout=5")
274 .header(TRANSFER_ENCODING, "chunked")
275 .header("content-type", "text/plain")
276 .header("trailer", "expires")
278 .header("proxy-authenticate", "Basic")
279 .body(())
280 .unwrap();
281
282 adapt_response_version(&mut resp, &ctx_with_version(Version::HTTP_2)).unwrap();
283
284 assert_eq!(resp.version(), Version::HTTP_2);
285 assert!(!resp.headers().contains_key(CONNECTION));
286 assert!(!resp.headers().contains_key("keep-alive"));
287 assert!(!resp.headers().contains_key(TRANSFER_ENCODING));
288 assert_eq!(resp.headers().get("content-type").unwrap(), "text/plain");
289 assert_eq!(resp.headers().get("trailer").unwrap(), "expires");
291 assert_eq!(resp.headers().get("proxy-authenticate").unwrap(), "Basic");
292 }
293
294 #[test]
295 fn test_h1_to_h2_non_websocket_101_errors() {
296 let mut resp = Response::builder()
297 .version(Version::HTTP_11)
298 .status(StatusCode::SWITCHING_PROTOCOLS)
299 .header(UPGRADE, "h2c")
300 .header(CONNECTION, "Upgrade")
301 .body(())
302 .unwrap();
303
304 let err =
307 adapt_response_version(&mut resp, &ctx_with_version(Version::HTTP_2)).unwrap_err();
308 assert!(
309 err.to_string().contains("only websocket is supported"),
310 "{err}"
311 );
312 }
313
314 #[test]
315 fn test_h2_to_h1_unsupported_extended_connect_errors() {
316 let mut resp = Response::builder()
317 .version(Version::HTTP_2)
318 .status(StatusCode::OK)
319 .body(())
320 .unwrap();
321
322 let err =
325 adapt_response_version(&mut resp, &connect_udp_ctx(Version::HTTP_11)).unwrap_err();
326 assert!(
327 err.to_string().contains("only websocket is supported"),
328 "{err}"
329 );
330 }
331
332 #[test]
333 fn test_h1_to_h2_websocket_101_becomes_200() {
334 let mut resp = Response::builder()
335 .version(Version::HTTP_11)
336 .status(StatusCode::SWITCHING_PROTOCOLS)
337 .header(UPGRADE, "websocket")
338 .header(CONNECTION, "Upgrade")
339 .header("sec-websocket-accept", SAMPLE_ACCEPT)
340 .body(())
341 .unwrap();
342
343 adapt_response_version(&mut resp, &websocket_ctx(Version::HTTP_2)).unwrap();
344
345 assert_eq!(resp.status(), StatusCode::OK);
346 assert_eq!(resp.version(), Version::HTTP_2);
347 assert!(!resp.headers().contains_key(UPGRADE));
349 assert!(!resp.headers().contains_key(CONNECTION));
350 assert!(!resp.headers().contains_key("sec-websocket-accept"));
351 }
352
353 #[test]
354 fn test_h2_to_h1_websocket_200_becomes_101() {
355 let mut resp = Response::builder()
356 .version(Version::HTTP_2)
357 .status(StatusCode::OK)
358 .body(())
359 .unwrap();
360
361 adapt_response_version(&mut resp, &websocket_ctx(Version::HTTP_11)).unwrap();
362
363 assert_eq!(resp.version(), Version::HTTP_11);
364 assert_eq!(resp.status(), StatusCode::SWITCHING_PROTOCOLS);
365 assert!(
366 resp.headers()
367 .typed_get::<Upgrade>()
368 .is_some_and(|u| u.is_websocket())
369 );
370 assert!(
371 resp.headers()
372 .typed_get::<Connection>()
373 .is_some_and(|c| c.contains_upgrade())
374 );
375 assert_eq!(
377 resp.headers().get("sec-websocket-accept").unwrap(),
378 SAMPLE_ACCEPT,
379 );
380 }
381
382 #[test]
383 fn test_h2_to_h1_non_websocket_only_changes_version() {
384 let mut resp = Response::builder()
385 .version(Version::HTTP_2)
386 .status(StatusCode::OK)
387 .header("content-type", "text/plain")
388 .body(())
389 .unwrap();
390
391 adapt_response_version(&mut resp, &ctx_with_version(Version::HTTP_11)).unwrap();
392
393 assert_eq!(resp.version(), Version::HTTP_11);
394 assert_eq!(resp.status(), StatusCode::OK);
395 assert_eq!(resp.headers().get("content-type").unwrap(), "text/plain");
396 }
397}