tuwunel_service/pusher/
request.rs1use std::{fmt::Debug, mem::swap};
2
3use bytes::BytesMut;
4use http::Response as HttpResponse;
5use reqwest::{Error as ReqwestError, Response};
6use ruma::api::{
7 IncomingResponse, OutgoingRequest, OutgoingRequestExt, auth_scheme::AuthScheme,
8 path_builder::PathBuilder,
9};
10use tuwunel_core::{
11 Err, Result, debug_warn, err, error::error_chain, implement, trace, utils::string_from_bytes,
12 warn,
13};
14use url::Url;
15
16use crate::client::read_response_capped;
17
18#[implement(super::Service)]
19#[tracing::instrument(level = "debug", skip_all)]
20pub(super) async fn send_request<T>(&self, dest: &str, request: T) -> Result<T::IncomingResponse>
21where
22 T: OutgoingRequest + Debug + Send,
23 for<'a> T::Authentication: AuthScheme<Input<'a> = ()>,
24 for<'a> T::PathBuilder: PathBuilder<Input<'a> = ()>,
25{
26 let dest = if dest.contains(['?', '#']) {
27 let parsed = Url::parse(dest).ok();
28
29 warn!(
30 gateway_host = parsed
31 .as_ref()
32 .and_then(|url| url.host_str())
33 .unwrap_or("<invalid>"),
34 has_query = dest.contains('?'),
35 has_fragment = dest.contains('#'),
36 "Push gateway URL carries a query string or fragment, which is not supported; the \
37 notification path is appended after it",
38 );
39
40 dest
41 } else {
42 let push_path = self
43 .services
44 .config
45 .notification_push_path
46 .trim_end_matches('/');
47
48 let dest = dest.trim_end_matches('/');
49
50 dest.strip_suffix(push_path).unwrap_or(dest)
51 };
52
53 trace!("Push gateway destination: {dest}");
54
55 let http_request = request
56 .try_into_http_request::<BytesMut>(dest, (), ())
57 .map_err(|e| {
58 err!(BadServerResponse(warn!(
59 "Failed to find destination {dest} for push gateway: {e}"
60 )))
61 })?
62 .map(BytesMut::freeze);
63
64 let reqwest_request = reqwest::Request::try_from(http_request)?;
65
66 if self
67 .services
68 .client
69 .proxy
70 .resolver_alias(reqwest_request.url())
71 {
72 return Err!(BadServerResponse(
73 "Not allowed to request a locally resolved proxy endpoint"
74 ));
75 }
76
77 trace!("Checking request URL for IP");
78 if !self
79 .services
80 .client
81 .valid_cidr_range_url(reqwest_request.url())
82 {
83 return Err!(BadServerResponse("Not allowed to send requests to this IP"));
84 }
85
86 match self
87 .services
88 .client
89 .pusher
90 .execute(reqwest_request)
91 .await
92 {
93 | Err(error) => handler_err(dest, error),
94 | Ok(response) => self.handle_ok::<T>(dest, response).await,
95 }
96}
97
98#[implement(super::Service)]
99async fn handle_ok<T>(&self, dest: &str, mut response: Response) -> Result<T::IncomingResponse>
100where
101 T: OutgoingRequest,
102{
103 trace!("Checking response destination's IP");
104 if let Some(remote_addr) = response.remote_addr()
105 && !self
106 .services
107 .client
108 .valid_cidr_range_ip(remote_addr.ip())
109 && !self.services.client.proxied(response.url())
110 {
111 return Err!(BadServerResponse("Not allowed to send requests to this IP"));
112 }
113
114 let status = response.status();
115 let mut http_response_builder = HttpResponse::builder()
116 .status(status)
117 .version(response.version());
118
119 swap(
120 response.headers_mut(),
121 http_response_builder
122 .headers_mut()
123 .expect("http::response::Builder is usable"),
124 );
125
126 let limit = self.services.config.max_response_size;
127 let body = read_response_capped(response, limit).await?;
128
129 if !status.is_success() {
130 debug_warn!(body = ?string_from_bytes(&body), "Push gateway response");
131 return Err!(BadServerResponse(warn!(
132 "Push gateway {dest} returned unsuccessful HTTP response: {status}"
133 )));
134 }
135
136 let response = T::IncomingResponse::try_from_http_response(
137 http_response_builder
138 .body(body)
139 .expect("reqwest body is valid http body"),
140 );
141
142 response.map_err(|e| {
143 err!(BadServerResponse(warn!("Push gateway {dest} returned invalid response: {e}")))
144 })
145}
146
147fn handler_err<R>(dest: &str, error: ReqwestError) -> Result<R> {
148 warn!(%dest, chain = %error_chain(&error), "Could not send request to pusher");
149 Err(error.into())
150}