1pub mod proxy;
18pub mod reload;
19
20use ferryman_edge_core::{check, Claims, JwtVerifier, Limiter, ReloadingTls, SharedTable};
21use http::{HeaderValue, Request, Response};
22use hyper::body::{Bytes, Incoming};
23use hyper::service::service_fn;
24use hyper_util::client::legacy::connect::HttpConnector;
25use hyper_util::client::legacy::Client;
26use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
27use hyper_util::server::conn::auto;
28use hyper_util::server::graceful::{GracefulShutdown, Watcher};
29use std::convert::Infallible;
30use std::future::Future;
31use std::net::SocketAddr;
32use std::sync::atomic::{AtomicBool, Ordering};
33use std::sync::Arc;
34use std::time::{Duration, Instant};
35use tokio::net::{TcpListener, TcpStream};
36use tokio_rustls::TlsAcceptor;
37
38pub type UpstreamClient = Client<HttpConnector, proxy::Body>;
41
42pub struct AppState {
44 pub tls: Arc<ReloadingTls>,
45 pub table: SharedTable,
46 pub jwt: Arc<JwtVerifier>,
47 pub limiter: Option<Arc<Limiter>>,
48 pub client: UpstreamClient,
49}
50
51const TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
52const HEADER_READ_TIMEOUT: Duration = Duration::from_secs(10);
53const SHUTDOWN_DRAIN: Duration = Duration::from_secs(25);
54const ACCEPT_ERROR_BACKOFF: Duration = Duration::from_millis(50);
55const H2_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30);
56const H2_MAX_CONCURRENT_STREAMS: u32 = 64;
61
62pub async fn serve(
66 listener: TcpListener,
67 state: Arc<AppState>,
68 shutdown: impl Future<Output = ()>,
69) {
70 let graceful = GracefulShutdown::new();
71 let mut http = auto::Builder::new(TokioExecutor::new());
72 http.http1()
73 .timer(TokioTimer::new())
74 .header_read_timeout(HEADER_READ_TIMEOUT);
75 http.http2()
76 .timer(TokioTimer::new())
77 .keep_alive_interval(H2_KEEPALIVE_INTERVAL)
78 .max_concurrent_streams(H2_MAX_CONCURRENT_STREAMS);
79
80 tokio::pin!(shutdown);
81 loop {
82 tokio::select! {
83 _ = &mut shutdown => break,
84 accepted = listener.accept() => {
85 match accepted {
86 Ok((stream, peer)) => {
87 let state = state.clone();
88 let http = http.clone();
89 let watcher = graceful.watcher();
90 tokio::spawn(async move {
91 handle_conn(stream, peer, state, http, watcher).await;
92 });
93 }
94 Err(e) => {
95 tracing::warn!(?e, "accept error");
97 tokio::time::sleep(ACCEPT_ERROR_BACKOFF).await;
98 }
99 }
100 }
101 }
102 }
103
104 drop(listener);
105 tracing::info!("shutdown signal received; draining connections");
106 tokio::select! {
107 _ = graceful.shutdown() => tracing::info!("all connections drained"),
108 _ = tokio::time::sleep(SHUTDOWN_DRAIN) => {
109 tracing::warn!("graceful shutdown timed out; dropping remaining connections");
110 }
111 }
112}
113
114async fn handle_conn(
115 stream: TcpStream,
116 peer: SocketAddr,
117 state: Arc<AppState>,
118 http: auto::Builder<TokioExecutor>,
119 watcher: Watcher,
120) {
121 let acceptor = TlsAcceptor::from(state.tls.current());
122 let started = Instant::now();
123 let tls_stream =
124 match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, acceptor.accept(stream)).await {
125 Ok(Ok(s)) => s,
126 Ok(Err(e)) => {
127 metrics::counter!("ferryman_tls_handshake_failures_total").increment(1);
128 tracing::debug!(?peer, ?e, "tls handshake failed");
129 return;
130 }
131 Err(_) => {
132 metrics::counter!("ferryman_tls_handshake_failures_total").increment(1);
133 tracing::debug!(?peer, "tls handshake timed out");
134 return;
135 }
136 };
137 metrics::histogram!("ferryman_tls_handshake_seconds").record(started.elapsed().as_secs_f64());
138
139 let http = if tls_stream.get_ref().1.alpn_protocol() == Some(b"h2") {
143 http.http2_only()
144 } else {
145 http.http1_only()
146 };
147 let io = TokioIo::new(tls_stream);
148 let seen_request = Arc::new(AtomicBool::new(false));
149 let seen = seen_request.clone();
150 let svc = service_fn(move |req| {
151 seen.store(true, Ordering::Release);
152 let state = state.clone();
153 async move { Ok::<_, Infallible>(route_request(state, req, peer).await) }
154 });
155
156 let conn = watcher.watch(http.serve_connection(io, svc).into_owned());
157 let first_request = tokio::time::sleep(HEADER_READ_TIMEOUT);
162 tokio::pin!(conn, first_request);
163 let result = tokio::select! {
164 r = &mut conn => r,
165 _ = &mut first_request => {
166 if !seen_request.load(Ordering::Acquire) {
167 tracing::debug!(?peer, "no request before timeout; closing");
168 return;
169 }
170 conn.await
171 }
172 };
173 if let Err(e) = result {
174 tracing::debug!(?peer, ?e, "connection error");
175 }
176}
177
178async fn route_request(
180 state: Arc<AppState>,
181 mut req: Request<Incoming>,
182 peer: SocketAddr,
183) -> Response<proxy::Body> {
184 let claims = match authenticate(&state.jwt, &req).await {
185 Ok(c) => c,
186 Err(reason) => {
187 metrics::counter!("ferryman_requests_total", "status" => "401").increment(1);
188 metrics::counter!("ferryman_auth_failures_total", "reason" => reason.as_str())
189 .increment(1);
190 return unauthorized_response();
191 }
192 };
193
194 if let Some(limiter) = &state.limiter {
195 if !check(limiter, &claims.sub) {
196 metrics::counter!("ferryman_requests_total", "status" => "429").increment(1);
197 metrics::counter!("ferryman_ratelimited_total").increment(1);
199 return rate_limited_response();
200 }
201 }
202
203 proxy::strip_hop_by_hop(req.headers_mut());
207 req.headers_mut().remove("x-ferryman-tenant");
210 if let Ok(v) = HeaderValue::from_str(&claims.sub) {
211 req.headers_mut().insert("x-ferryman-tenant", v);
212 }
213
214 match proxy::handle(state.table.clone(), state.client.clone(), req, peer.ip()).await {
215 Ok(resp) => resp,
216 Err(e) => {
217 tracing::error!(?e, "unhandled proxy error");
218 metrics::counter!("ferryman_requests_total", "status" => "500").increment(1);
219 internal_error_response()
220 }
221 }
222}
223
224enum AuthFailure {
225 Missing,
226 Invalid,
227}
228
229impl AuthFailure {
230 fn as_str(&self) -> &'static str {
231 match self {
232 AuthFailure::Missing => "missing",
233 AuthFailure::Invalid => "invalid",
234 }
235 }
236}
237
238async fn authenticate(jwt: &JwtVerifier, req: &Request<Incoming>) -> Result<Claims, AuthFailure> {
239 let raw = req
240 .headers()
241 .get(http::header::AUTHORIZATION)
242 .ok_or(AuthFailure::Missing)?;
243 let raw = raw.to_str().map_err(|_| AuthFailure::Invalid)?;
244 let (scheme, token) = raw.split_once(' ').ok_or(AuthFailure::Invalid)?;
245 if !scheme.eq_ignore_ascii_case("bearer") || token.is_empty() {
246 return Err(AuthFailure::Invalid);
247 }
248 jwt.verify(token).await.ok_or(AuthFailure::Invalid)
249}
250
251fn unauthorized_response() -> Response<proxy::Body> {
252 Response::builder()
253 .status(401)
254 .header(http::header::WWW_AUTHENTICATE, "Bearer")
255 .body(proxy::text_body(Bytes::from_static(b"unauthorized")))
256 .expect("static response builds")
257}
258
259fn rate_limited_response() -> Response<proxy::Body> {
260 Response::builder()
261 .status(429)
262 .header("retry-after", "1")
263 .body(proxy::text_body(Bytes::from_static(b"rate limited")))
264 .expect("static response builds")
265}
266
267fn internal_error_response() -> Response<proxy::Body> {
268 Response::builder()
269 .status(500)
270 .body(proxy::text_body(Bytes::from_static(b"internal error")))
271 .expect("static response builds")
272}