use ferryman_edge_core::SharedTable;
use http::{HeaderMap, HeaderValue};
use http_body_util::{BodyExt, Full, LengthLimitError, Limited};
use hyper::body::Bytes;
use hyper::body::Incoming;
use hyper::{Request, Response};
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::client::legacy::Client;
use std::net::IpAddr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "boxed_body")]
pub type BoxErr = Box<dyn std::error::Error + Send + Sync>;
#[cfg(not(feature = "boxed_body"))]
pub type Body = Full<Bytes>;
#[cfg(feature = "boxed_body")]
pub type Body = http_body_util::combinators::BoxBody<Bytes, BoxErr>;
const MAX_BODY_BYTES: usize = 8 * 1024 * 1024;
const UPSTREAM_TIMEOUT: Duration = Duration::from_secs(30);
const BODY_READ_TIMEOUT: Duration = Duration::from_secs(30);
const HOP_BY_HOP_HEADERS: &[&str] = &[
"connection",
"keep-alive",
"proxy-connection",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"proxy-authenticate",
"proxy-authorization",
];
pub(crate) fn strip_hop_by_hop(headers: &mut HeaderMap) {
let mut extra: Vec<String> = Vec::new();
for v in headers.get_all(http::header::CONNECTION) {
if let Ok(s) = v.to_str() {
extra.extend(
s.split(',')
.map(|p| p.trim().to_ascii_lowercase())
.filter(|p| !p.is_empty()),
);
}
}
for name in HOP_BY_HOP_HEADERS {
headers.remove(*name);
}
for name in extra {
headers.remove(name.as_str());
}
}
fn set_forwarded(headers: &mut HeaderMap, ip: IpAddr) {
headers.remove("forwarded");
headers.remove("x-real-ip");
headers.insert(
"x-forwarded-for",
HeaderValue::from_str(&ip.to_string()).expect("an IP is a valid header value"),
);
headers.insert("x-forwarded-proto", HeaderValue::from_static("https"));
}
#[cfg(not(feature = "boxed_body"))]
pub(crate) fn text_body(bytes: Bytes) -> Body {
Full::new(bytes)
}
#[cfg(feature = "boxed_body")]
pub(crate) fn text_body(bytes: Bytes) -> Body {
Full::new(bytes)
.map_err(|never: std::convert::Infallible| -> BoxErr { match never {} })
.boxed()
}
fn plain(status: u16, msg: &'static [u8]) -> anyhow::Result<Response<Body>> {
metrics::counter!("ferryman_requests_total", "status" => status.to_string()).increment(1);
Ok(Response::builder()
.status(status)
.body(text_body(Bytes::from_static(msg)))?)
}
fn has_dot_segment(path: &str) -> bool {
path.split('/').any(|seg| {
let seg = seg.to_ascii_lowercase().replace("%2e", ".");
seg == "." || seg == ".."
})
}
fn is_client_body_error(e: &(dyn std::error::Error + 'static)) -> bool {
error_chain(e).any(|c| {
c.is::<LengthLimitError>()
|| c.downcast_ref::<hyper::Error>()
.is_some_and(|h| h.is_user())
})
}
fn error_chain<'a>(
e: &'a (dyn std::error::Error + 'static),
) -> impl Iterator<Item = &'a (dyn std::error::Error + 'static)> {
std::iter::successors(Some(e), |c| c.source())
}
pub async fn handle(
table: SharedTable,
client: Client<HttpConnector, Body>,
req: Request<Incoming>,
peer_ip: IpAddr,
) -> Result<Response<Body>, anyhow::Error> {
let started = std::time::Instant::now();
let snapshot = table.load();
let path = req.uri().path().to_string();
if has_dot_segment(&path) {
return plain(400, b"bad path");
}
if req
.headers()
.get(http::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
.is_some_and(|len| len > MAX_BODY_BYTES as u64)
{
return plain(413, b"payload too large");
}
let (mut parts, body) = req.into_parts();
let (fwd_body, upload_done) =
match tokio::time::timeout(BODY_READ_TIMEOUT, forward_body(body)).await {
Ok(Ok(b)) => b,
Ok(Err(e)) if e.downcast_ref::<LengthLimitError>().is_some() => {
return plain(413, b"payload too large");
}
Ok(Err(_)) => return plain(400, b"request body error"),
Err(_) => return plain(408, b"request body timeout"),
};
let upstream = match snapshot.lookup(&path) {
Some(u) => u.clone(),
None if snapshot.has_prefix(&path) => return plain(503, b"upstream unavailable"),
None => return plain(404, b"no route"),
};
let mut up_parts = upstream.uri.clone().into_parts();
up_parts.path_and_query = parts.uri.path_and_query().cloned();
parts.uri = http::Uri::from_parts(up_parts)?;
parts.version = http::Version::HTTP_11;
if let Some(authority) = upstream.uri.authority() {
parts.headers.insert(
http::header::HOST,
HeaderValue::from_str(authority.as_str())?,
);
}
set_forwarded(&mut parts.headers, peer_ip);
let deadline = tokio::time::Instant::now() + UPSTREAM_TIMEOUT;
let fwd = Request::from_parts(parts, fwd_body);
let host = upstream
.uri
.authority()
.map_or_else(String::new, |a| a.to_string());
let resp = match tokio::time::timeout_at(deadline, client.request(fwd)).await {
Ok(Ok(resp)) => resp,
Ok(Err(e)) if is_client_body_error(&e) => {
let too_large = error_chain(&e).any(|c| c.is::<LengthLimitError>());
return if too_large {
plain(413, b"payload too large")
} else {
plain(400, b"request body error")
};
}
Ok(Err(e)) => {
tracing::warn!(upstream = %host, error = %e, "upstream request failed");
upstream.mark_failed();
metrics::counter!("ferryman_requests_total", "status" => "502", "upstream" => host)
.increment(1);
return Ok(Response::builder()
.status(502)
.body(text_body(Bytes::from_static(b"bad gateway")))?);
}
Err(_) => {
if upload_done.load(Ordering::Acquire) {
upstream.mark_failed();
}
metrics::counter!("ferryman_requests_total", "status" => "504", "upstream" => host)
.increment(1);
return Ok(Response::builder()
.status(504)
.body(text_body(Bytes::from_static(b"upstream timeout")))?);
}
};
let status = resp.status();
let (mut resp_parts, resp_body) = resp.into_parts();
strip_hop_by_hop(&mut resp_parts.headers);
resp_parts.version = http::Version::default();
#[cfg(not(feature = "boxed_body"))]
let out_body: Body = match tokio::time::timeout_at(deadline, resp_body.collect()).await {
Ok(Ok(collected)) => Full::new(collected.to_bytes()),
Ok(Err(_)) | Err(_) => {
upstream.mark_failed();
metrics::counter!("ferryman_requests_total", "status" => "502", "upstream" => host)
.increment(1);
return Ok(Response::builder()
.status(502)
.body(text_body(Bytes::from_static(b"bad gateway")))?);
}
};
#[cfg(feature = "boxed_body")]
let out_body: Body = resp_body.map_err(Into::into).boxed();
if matches!(status.as_u16(), 502..=504) {
upstream.mark_failed();
} else {
upstream.mark_success();
}
metrics::histogram!("ferryman_request_duration_seconds", "upstream" => host.clone())
.record(started.elapsed().as_secs_f64());
metrics::counter!(
"ferryman_requests_total",
"status" => status.as_u16().to_string(),
"upstream" => host
)
.increment(1);
Ok(Response::from_parts(resp_parts, out_body))
}
type UploadDone = Arc<AtomicBool>;
#[cfg(not(feature = "boxed_body"))]
async fn forward_body(body: Incoming) -> anyhow::Result<(Body, UploadDone)> {
match Limited::new(body, MAX_BODY_BYTES).collect().await {
Ok(collected) => Ok((
Full::new(collected.to_bytes()),
Arc::new(AtomicBool::new(true)),
)),
Err(e) => match e.downcast::<LengthLimitError>() {
Ok(too_large) => Err(anyhow::Error::new(*too_large)),
Err(other) => Err(anyhow::anyhow!("{other}")),
},
}
}
#[cfg(feature = "boxed_body")]
async fn forward_body(body: Incoming) -> anyhow::Result<(Body, UploadDone)> {
let inner = Limited::new(body, MAX_BODY_BYTES);
let done: UploadDone = Arc::new(AtomicBool::new(hyper::body::Body::is_end_stream(&inner)));
let body = TrackEnd {
inner,
done: done.clone(),
};
Ok((body.boxed(), done))
}
#[cfg(feature = "boxed_body")]
struct TrackEnd<B> {
inner: B,
done: UploadDone,
}
#[cfg(feature = "boxed_body")]
impl<B: hyper::body::Body + Unpin> hyper::body::Body for TrackEnd<B> {
type Data = B::Data;
type Error = B::Error;
fn poll_frame(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Result<hyper::body::Frame<Self::Data>, Self::Error>>> {
let polled = std::pin::Pin::new(&mut self.inner).poll_frame(cx);
if matches!(polled, std::task::Poll::Ready(None)) || self.inner.is_end_stream() {
self.done.store(true, Ordering::Release);
}
polled
}
fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}
fn size_hint(&self) -> hyper::body::SizeHint {
self.inner.size_hint()
}
}
#[cfg(test)]
mod tests {
use super::has_dot_segment;
#[test]
fn dot_segments_are_detected() {
for bad in [
"/svc-a/../svc-b",
"/svc-a/./x",
"/svc-a/%2e%2e/svc-b",
"/svc-a/%2E/x",
"/..",
] {
assert!(has_dot_segment(bad), "{bad}");
}
for ok in ["/svc-a/x", "/svc-a/.hidden", "/svc-a/a..b", "/"] {
assert!(!has_dot_segment(ok), "{ok}");
}
}
}