1use ferryman_edge_core::SharedTable;
20use http::{HeaderMap, HeaderValue};
21use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
22use hyper::body::Bytes;
23use hyper::body::Incoming;
24use hyper::{Request, Response};
25use hyper_util::client::legacy::connect::HttpConnector;
26use hyper_util::client::legacy::Client;
27use std::net::IpAddr;
28use std::sync::atomic::{AtomicBool, Ordering};
29use std::sync::Arc;
30use std::time::Duration;
31
32#[cfg(feature = "boxed_body")]
33pub type BoxErr = Box<dyn std::error::Error + Send + Sync>;
34
35#[cfg(not(feature = "boxed_body"))]
36pub type Body = Full<Bytes>;
37#[cfg(feature = "boxed_body")]
38pub type Body = http_body_util::combinators::BoxBody<Bytes, BoxErr>;
39
40const MAX_BODY_BYTES: usize = 8 * 1024 * 1024;
44
45const UPSTREAM_TIMEOUT: Duration = Duration::from_secs(30);
48
49const BODY_READ_TIMEOUT: Duration = Duration::from_secs(30);
53
54const HOP_BY_HOP_HEADERS: &[&str] = &[
55 "connection",
56 "keep-alive",
57 "proxy-connection",
58 "te",
59 "trailer",
60 "transfer-encoding",
61 "upgrade",
62 "proxy-authenticate",
63 "proxy-authorization",
64];
65
66pub(crate) fn strip_hop_by_hop(headers: &mut HeaderMap) {
70 let mut extra: Vec<String> = Vec::new();
71 for v in headers.get_all(http::header::CONNECTION) {
72 if let Ok(s) = v.to_str() {
73 extra.extend(
74 s.split(',')
75 .map(|p| p.trim().to_ascii_lowercase())
76 .filter(|p| !p.is_empty()),
77 );
78 }
79 }
80 for name in HOP_BY_HOP_HEADERS {
81 headers.remove(*name);
82 }
83 for name in extra {
84 headers.remove(name.as_str());
85 }
86}
87
88fn set_forwarded(headers: &mut HeaderMap, ip: IpAddr) {
91 headers.remove("forwarded");
92 headers.remove("x-real-ip");
93 headers.insert(
94 "x-forwarded-for",
95 HeaderValue::from_str(&ip.to_string()).expect("an IP is a valid header value"),
96 );
97 headers.insert("x-forwarded-proto", HeaderValue::from_static("https"));
98}
99
100#[cfg(not(feature = "boxed_body"))]
101pub(crate) fn text_body(bytes: Bytes) -> Body {
102 Full::new(bytes)
103}
104
105#[cfg(feature = "boxed_body")]
106pub(crate) fn text_body(bytes: Bytes) -> Body {
107 Full::new(bytes)
108 .map_err(|never: std::convert::Infallible| -> BoxErr { match never {} })
109 .boxed()
110}
111
112fn plain(status: u16, msg: &'static [u8]) -> anyhow::Result<Response<Body>> {
113 metrics::counter!("ferryman_requests_total", "status" => status.to_string()).increment(1);
114 Ok(Response::builder()
115 .status(status)
116 .body(text_body(Bytes::from_static(msg)))?)
117}
118
119fn has_dot_segment(path: &str) -> bool {
123 path.split('/').any(|seg| {
124 let seg = seg.to_ascii_lowercase().replace("%2e", ".");
125 seg == "." || seg == ".."
126 })
127}
128
129fn is_client_body_error(e: &(dyn std::error::Error + 'static)) -> bool {
135 error_chain(e).any(|c| {
136 c.is::<LengthLimitError>()
137 || c.downcast_ref::<hyper::Error>()
138 .is_some_and(|h| h.is_user())
139 })
140}
141
142fn error_chain<'a>(
143 e: &'a (dyn std::error::Error + 'static),
144) -> impl Iterator<Item = &'a (dyn std::error::Error + 'static)> {
145 std::iter::successors(Some(e), |c| c.source())
146}
147
148pub async fn handle(
151 table: SharedTable,
152 client: Client<HttpConnector, Body>,
153 req: Request<Incoming>,
154 peer_ip: IpAddr,
155) -> Result<Response<Body>, anyhow::Error> {
156 let started = std::time::Instant::now();
157 let snapshot = table.load();
158 let path = req.uri().path().to_string();
159
160 if has_dot_segment(&path) {
161 return plain(400, b"bad path");
162 }
163
164 if req
167 .headers()
168 .get(http::header::CONTENT_LENGTH)
169 .and_then(|v| v.to_str().ok())
170 .and_then(|s| s.parse::<u64>().ok())
171 .is_some_and(|len| len > MAX_BODY_BYTES as u64)
172 {
173 return plain(413, b"payload too large");
174 }
175
176 let (mut parts, body) = req.into_parts();
182 let (fwd_body, upload_done) =
183 match tokio::time::timeout(BODY_READ_TIMEOUT, forward_body(body)).await {
184 Ok(Ok(b)) => b,
185 Ok(Err(e)) if e.downcast_ref::<LengthLimitError>().is_some() => {
186 return plain(413, b"payload too large");
187 }
188 Ok(Err(_)) => return plain(400, b"request body error"),
189 Err(_) => return plain(408, b"request body timeout"),
190 };
191
192 let upstream = match snapshot.lookup(&path) {
193 Some(u) => u.clone(),
194 None if snapshot.has_prefix(&path) => return plain(503, b"upstream unavailable"),
196 None => return plain(404, b"no route"),
197 };
198
199 let mut up_parts = upstream.uri.clone().into_parts();
203 up_parts.path_and_query = parts.uri.path_and_query().cloned();
204 parts.uri = http::Uri::from_parts(up_parts)?;
205 parts.version = http::Version::HTTP_11;
209 if let Some(authority) = upstream.uri.authority() {
210 parts.headers.insert(
211 http::header::HOST,
212 HeaderValue::from_str(authority.as_str())?,
213 );
214 }
215 set_forwarded(&mut parts.headers, peer_ip);
216
217 let deadline = tokio::time::Instant::now() + UPSTREAM_TIMEOUT;
220 let fwd = Request::from_parts(parts, fwd_body);
221
222 let host = upstream
224 .uri
225 .authority()
226 .map_or_else(String::new, |a| a.to_string());
227 let resp = match tokio::time::timeout_at(deadline, client.request(fwd)).await {
228 Ok(Ok(resp)) => resp,
229 Ok(Err(e)) if is_client_body_error(&e) => {
230 let too_large = error_chain(&e).any(|c| c.is::<LengthLimitError>());
231 return if too_large {
232 plain(413, b"payload too large")
233 } else {
234 plain(400, b"request body error")
235 };
236 }
237 Ok(Err(e)) => {
238 tracing::warn!(upstream = %host, error = %e, "upstream request failed");
239 upstream.mark_failed();
240 metrics::counter!("ferryman_requests_total", "status" => "502", "upstream" => host)
241 .increment(1);
242 return Ok(Response::builder()
243 .status(502)
244 .body(text_body(Bytes::from_static(b"bad gateway")))?);
245 }
246 Err(_) => {
247 if upload_done.load(Ordering::Acquire) {
250 upstream.mark_failed();
251 }
252 metrics::counter!("ferryman_requests_total", "status" => "504", "upstream" => host)
253 .increment(1);
254 return Ok(Response::builder()
255 .status(504)
256 .body(text_body(Bytes::from_static(b"upstream timeout")))?);
257 }
258 };
259
260 let status = resp.status();
261 let (mut resp_parts, resp_body) = resp.into_parts();
262 strip_hop_by_hop(&mut resp_parts.headers);
263 resp_parts.version = http::Version::default();
266
267 #[cfg(not(feature = "boxed_body"))]
268 let out_body: Body = match tokio::time::timeout_at(deadline, resp_body.collect()).await {
269 Ok(Ok(collected)) => Full::new(collected.to_bytes()),
270 Ok(Err(_)) | Err(_) => {
271 upstream.mark_failed();
272 metrics::counter!("ferryman_requests_total", "status" => "502", "upstream" => host)
273 .increment(1);
274 return Ok(Response::builder()
275 .status(502)
276 .body(text_body(Bytes::from_static(b"bad gateway")))?);
277 }
278 };
279 #[cfg(feature = "boxed_body")]
280 let out_body: Body = resp_body.map_err(Into::into).boxed();
281
282 if matches!(status.as_u16(), 502..=504) {
286 upstream.mark_failed();
287 } else {
288 upstream.mark_success();
289 }
290 metrics::histogram!("ferryman_request_duration_seconds", "upstream" => host.clone())
291 .record(started.elapsed().as_secs_f64());
292 metrics::counter!(
293 "ferryman_requests_total",
294 "status" => status.as_u16().to_string(),
295 "upstream" => host
296 )
297 .increment(1);
298
299 Ok(Response::from_parts(resp_parts, out_body))
300}
301
302type UploadDone = Arc<AtomicBool>;
313
314#[cfg(not(feature = "boxed_body"))]
315async fn forward_body(body: Incoming) -> anyhow::Result<(Body, UploadDone)> {
316 match Limited::new(body, MAX_BODY_BYTES).collect().await {
317 Ok(collected) => Ok((
318 Full::new(collected.to_bytes()),
319 Arc::new(AtomicBool::new(true)),
320 )),
321 Err(e) => match e.downcast::<LengthLimitError>() {
322 Ok(too_large) => Err(anyhow::Error::new(*too_large)),
323 Err(other) => Err(anyhow::anyhow!("{other}")),
324 },
325 }
326}
327
328#[cfg(feature = "boxed_body")]
332async fn forward_body(body: Incoming) -> anyhow::Result<(Body, UploadDone)> {
333 let inner = Limited::new(body, MAX_BODY_BYTES);
334 let done: UploadDone = Arc::new(AtomicBool::new(hyper::body::Body::is_end_stream(&inner)));
337 let body = TrackEnd {
338 inner,
339 done: done.clone(),
340 };
341 Ok((body.boxed(), done))
342}
343
344#[cfg(feature = "boxed_body")]
346struct TrackEnd<B> {
347 inner: B,
348 done: UploadDone,
349}
350
351#[cfg(feature = "boxed_body")]
352impl<B: hyper::body::Body + Unpin> hyper::body::Body for TrackEnd<B> {
353 type Data = B::Data;
354 type Error = B::Error;
355
356 fn poll_frame(
357 mut self: std::pin::Pin<&mut Self>,
358 cx: &mut std::task::Context<'_>,
359 ) -> std::task::Poll<Option<Result<hyper::body::Frame<Self::Data>, Self::Error>>> {
360 let polled = std::pin::Pin::new(&mut self.inner).poll_frame(cx);
361 if matches!(polled, std::task::Poll::Ready(None)) || self.inner.is_end_stream() {
362 self.done.store(true, Ordering::Release);
363 }
364 polled
365 }
366
367 fn is_end_stream(&self) -> bool {
368 self.inner.is_end_stream()
369 }
370
371 fn size_hint(&self) -> hyper::body::SizeHint {
372 self.inner.size_hint()
373 }
374}
375
376#[cfg(test)]
377mod tests {
378 use super::has_dot_segment;
379
380 #[test]
381 fn dot_segments_are_detected() {
382 for bad in [
383 "/svc-a/../svc-b",
384 "/svc-a/./x",
385 "/svc-a/%2e%2e/svc-b",
386 "/svc-a/%2E/x",
387 "/..",
388 ] {
389 assert!(has_dot_segment(bad), "{bad}");
390 }
391 for ok in ["/svc-a/x", "/svc-a/.hidden", "/svc-a/a..b", "/"] {
392 assert!(!has_dot_segment(ok), "{ok}");
393 }
394 }
395}