1use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
8use std::sync::Arc;
9
10use http_body_util::Full;
11use hyper::body::{Bytes, Incoming};
12use hyper::server::conn::http1;
13use hyper::service::service_fn;
14use hyper::{Request, Response, StatusCode, header};
15use hyper_util::client::legacy::Client;
16use hyper_util::rt::{TokioExecutor, TokioIo};
17use tokio::io::copy_bidirectional;
18use tokio::net::TcpListener;
19
20use stackless_core::types::TcpPort;
21
22use crate::state::DaemonState;
23
24pub const DEFAULT_PROXY_PORT: u16 = 4444;
28
29pub fn proxy_port() -> TcpPort {
30 std::env::var("STACKLESS_PROXY_PORT")
31 .ok()
32 .and_then(|value| value.parse().ok())
33 .and_then(|raw| TcpPort::try_new(raw).ok())
34 .unwrap_or_else(|| TcpPort::from_os(DEFAULT_PROXY_PORT))
35}
36
37type ProxyClient = Client<hyper_util::client::legacy::connect::HttpConnector, Incoming>;
38
39pub async fn serve(state: Arc<DaemonState>, port: TcpPort) -> std::io::Result<()> {
40 let port = port.get();
41 let v4 = TcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, port))).await?;
42 let v6 = TcpListener::bind(SocketAddr::from((Ipv6Addr::LOCALHOST, port))).await?;
43 let client: ProxyClient = Client::builder(TokioExecutor::new())
44 .build(hyper_util::client::legacy::connect::HttpConnector::new());
45
46 let accept = |listener: TcpListener, state: Arc<DaemonState>, client: ProxyClient| async move {
47 loop {
48 let Ok((stream, _)) = listener.accept().await else {
49 continue;
50 };
51 let state = state.clone();
52 let client = client.clone();
53 tokio::spawn(async move {
54 let service = service_fn(move |req| handle(req, state.clone(), client.clone()));
55 let _ = http1::Builder::new()
56 .serve_connection(TokioIo::new(stream), service)
57 .with_upgrades()
58 .await;
59 });
60 }
61 };
62 tokio::join!(
63 accept(v4, state.clone(), client.clone()),
64 accept(v6, state, client)
65 );
66 Ok(())
67}
68
69async fn handle(
70 mut request: Request<Incoming>,
71 state: Arc<DaemonState>,
72 client: ProxyClient,
73) -> Result<Response<ProxyBody>, hyper::Error> {
74 let host = request
75 .headers()
76 .get(hyper::header::HOST)
77 .and_then(|value| value.to_str().ok())
78 .map(|value| value.split(':').next().unwrap_or(value).to_owned());
79 let Some(host) = host else {
80 return Ok(text_response(
81 StatusCode::BAD_REQUEST,
82 "stackless proxy: no Host header",
83 ));
84 };
85 let Some(port) = state.route_lookup(&host) else {
86 return Ok(text_response(
87 StatusCode::NOT_FOUND,
88 &format!(
89 "stackless proxy: no instance route for {host:?}; `stackless list` shows live instances"
90 ),
91 ));
92 };
93
94 let upgrade = is_websocket_upgrade(&request);
95 let client_upgrade = upgrade.then(|| hyper::upgrade::on(&mut request));
96 let (mut parts, body) = request.into_parts();
97 let path_and_query = parts
98 .uri
99 .path_and_query()
100 .map(|pq| pq.as_str())
101 .unwrap_or("/");
102 let upstream_uri = format!("http://127.0.0.1:{}{path_and_query}", port.get());
103 match upstream_uri.parse() {
104 Ok(uri) => parts.uri = uri,
105 Err(_) => {
106 return Ok(text_response(
107 StatusCode::BAD_REQUEST,
108 "stackless proxy: unparseable request target",
109 ));
110 }
111 }
112 match client.request(Request::from_parts(parts, body)).await {
113 Ok(mut response) => {
114 if upgrade {
115 if response.status() != StatusCode::SWITCHING_PROTOCOLS {
116 return Ok(text_response(
117 StatusCode::BAD_GATEWAY,
118 &format!(
119 "stackless proxy: upstream on port {} rejected websocket upgrade with {}",
120 port.get(),
121 response.status()
122 ),
123 ));
124 }
125 let Some(client_upgrade) = client_upgrade else {
126 return Ok(text_response(
127 StatusCode::BAD_GATEWAY,
128 "stackless proxy: client upgrade was unavailable",
129 ));
130 };
131 let upstream_upgrade = hyper::upgrade::on(&mut response);
132 tokio::spawn(async move {
133 let Ok(client) = client_upgrade.await else {
134 return;
135 };
136 let Ok(upstream) = upstream_upgrade.await else {
137 return;
138 };
139 let mut client = TokioIo::new(client);
140 let mut upstream = TokioIo::new(upstream);
141 let _ = copy_bidirectional(&mut client, &mut upstream).await;
142 });
143 }
144 Ok(response.map(ProxyBody::Upstream))
145 }
146 Err(err) => Ok(text_response(
147 StatusCode::BAD_GATEWAY,
148 &format!(
149 "stackless proxy: upstream on port {} refused: {err}",
150 port.get()
151 ),
152 )),
153 }
154}
155
156fn is_websocket_upgrade<B>(request: &Request<B>) -> bool {
157 header_contains_token(request.headers(), header::CONNECTION, "upgrade")
158 && header_contains_token(request.headers(), header::UPGRADE, "websocket")
159}
160
161fn header_contains_token(
162 headers: &header::HeaderMap,
163 name: header::HeaderName,
164 token: &str,
165) -> bool {
166 headers.get_all(name).iter().any(|value| {
167 value.to_str().is_ok_and(|value| {
168 value
169 .split(',')
170 .map(str::trim)
171 .any(|part| part.eq_ignore_ascii_case(token))
172 })
173 })
174}
175
176#[derive(Debug)]
178pub enum ProxyBody {
179 Upstream(Incoming),
180 Text(Full<Bytes>),
181}
182
183impl hyper::body::Body for ProxyBody {
184 type Data = Bytes;
185 type Error = hyper::Error;
186
187 fn poll_frame(
188 self: std::pin::Pin<&mut Self>,
189 cx: &mut std::task::Context<'_>,
190 ) -> std::task::Poll<Option<Result<hyper::body::Frame<Self::Data>, Self::Error>>> {
191 match self.get_mut() {
193 Self::Upstream(incoming) => std::pin::Pin::new(incoming).poll_frame(cx),
194 Self::Text(full) => std::pin::Pin::new(full)
195 .poll_frame(cx)
196 .map_err(|never| match never {}),
197 }
198 }
199}
200
201fn text_response(status: StatusCode, message: &str) -> Response<ProxyBody> {
202 let mut response = Response::new(ProxyBody::Text(Full::new(Bytes::from(format!(
203 "{message}\n"
204 )))));
205 *response.status_mut() = status;
206 response
207}
208
209#[cfg(test)]
210mod tests {
211 #![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
212
213 use std::future::Future;
214
215 use tokio::io::{AsyncReadExt, AsyncWriteExt};
216 use tokio::net::TcpStream;
217
218 use stackless_core::types::{ProxyHost, TcpPort};
219
220 use super::*;
221
222 async fn start_raw_upstream<F, Fut>(handler: F) -> (u16, tokio::task::JoinHandle<()>)
223 where
224 F: FnOnce(TcpStream) -> Fut + Send + 'static,
225 Fut: Future<Output = ()> + Send + 'static,
226 {
227 let listener = TcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0)))
228 .await
229 .unwrap();
230 let port = listener.local_addr().unwrap().port();
231 let handle = tokio::spawn(async move {
232 let (stream, _) = listener.accept().await.unwrap();
233 handler(stream).await;
234 });
235 (port, handle)
236 }
237
238 async fn start_proxy(state: Arc<DaemonState>) -> (u16, tokio::task::JoinHandle<()>) {
239 let std_listener =
240 std::net::TcpListener::bind(SocketAddr::from((Ipv4Addr::LOCALHOST, 0))).unwrap();
241 let port = std_listener.local_addr().unwrap().port();
242 drop(std_listener);
243 let handle = tokio::spawn(async move {
244 let _ = serve(state, TcpPort::from_os(port)).await;
245 });
246 for _ in 0..50 {
247 if TcpStream::connect(SocketAddr::from((Ipv4Addr::LOCALHOST, port)))
248 .await
249 .is_ok()
250 {
251 return (port, handle);
252 }
253 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
254 }
255 panic!("proxy did not start on port {port}");
256 }
257
258 async fn read_until_headers(stream: &mut TcpStream) -> Vec<u8> {
259 let mut bytes = Vec::new();
260 loop {
261 let mut buf = [0_u8; 256];
262 let n = stream.read(&mut buf).await.unwrap();
263 if n == 0 {
264 break;
265 }
266 bytes.extend_from_slice(&buf[..n]);
267 if bytes.windows(4).any(|window| window == b"\r\n\r\n") {
268 break;
269 }
270 }
271 bytes
272 }
273
274 #[tokio::test]
275 async fn normal_http_request_routes_to_upstream() {
276 let (upstream_port, upstream) = start_raw_upstream(|mut stream| async move {
277 let request = String::from_utf8(read_until_headers(&mut stream).await).unwrap();
278 assert!(request.starts_with("GET /hello?x=1 HTTP/1.1\r\n"));
279 stream
280 .write_all(b"HTTP/1.1 200 OK\r\ncontent-length: 2\r\n\r\nok")
281 .await
282 .unwrap();
283 })
284 .await;
285 let state = Arc::new(DaemonState::default());
286 state.route_set(
287 ProxyHost::try_new("demo.localhost").unwrap(),
288 TcpPort::from_os(upstream_port),
289 );
290 let (proxy_port, proxy) = start_proxy(state).await;
291
292 let mut client = TcpStream::connect(SocketAddr::from((Ipv4Addr::LOCALHOST, proxy_port)))
293 .await
294 .unwrap();
295 client
296 .write_all(
297 format!(
298 "GET /hello?x=1 HTTP/1.1\r\nHost: demo.localhost:{proxy_port}\r\nConnection: close\r\n\r\n"
299 )
300 .as_bytes(),
301 )
302 .await
303 .unwrap();
304 let mut response = String::new();
305 client.read_to_string(&mut response).await.unwrap();
306
307 assert!(response.starts_with("HTTP/1.1 200 OK\r\n"));
308 assert!(response.ends_with("\r\n\r\nok"));
309 upstream.await.unwrap();
310 proxy.abort();
311 }
312
313 #[tokio::test]
314 async fn missing_route_returns_stackless_404() {
315 let state = Arc::new(DaemonState::default());
316 let (proxy_port, proxy) = start_proxy(state).await;
317 let mut client = TcpStream::connect(SocketAddr::from((Ipv4Addr::LOCALHOST, proxy_port)))
318 .await
319 .unwrap();
320 client
321 .write_all(
322 format!(
323 "GET / HTTP/1.1\r\nHost: absent.localhost:{proxy_port}\r\nConnection: close\r\n\r\n"
324 )
325 .as_bytes(),
326 )
327 .await
328 .unwrap();
329 let mut response = String::new();
330 client.read_to_string(&mut response).await.unwrap();
331
332 assert!(response.starts_with("HTTP/1.1 404 Not Found\r\n"));
333 assert!(response.contains("stackless proxy: no instance route"));
334 proxy.abort();
335 }
336
337 #[tokio::test]
338 async fn websocket_upgrade_tunnels_bidirectional_bytes() {
339 let (upstream_port, upstream) = start_raw_upstream(|mut stream| async move {
340 let request = String::from_utf8(read_until_headers(&mut stream).await).unwrap();
341 assert!(request.starts_with("GET /hmr HTTP/1.1\r\n"));
342 let request_lower = request.to_ascii_lowercase();
343 assert!(request_lower.contains("connection: keep-alive, upgrade\r\n"));
344 assert!(request_lower.contains("upgrade: websocket\r\n"));
345 stream
346 .write_all(
347 b"HTTP/1.1 101 Switching Protocols\r\n\
348 Connection: Upgrade\r\n\
349 Upgrade: websocket\r\n\
350 \r\n",
351 )
352 .await
353 .unwrap();
354 let mut ping = [0_u8; 4];
355 stream.read_exact(&mut ping).await.unwrap();
356 assert_eq!(&ping, b"ping");
357 stream.write_all(b"pong").await.unwrap();
358 })
359 .await;
360 let state = Arc::new(DaemonState::default());
361 state.route_set(
362 ProxyHost::try_new("demo.localhost").unwrap(),
363 TcpPort::from_os(upstream_port),
364 );
365 let (proxy_port, proxy) = start_proxy(state).await;
366
367 let mut client = TcpStream::connect(SocketAddr::from((Ipv4Addr::LOCALHOST, proxy_port)))
368 .await
369 .unwrap();
370 client
371 .write_all(
372 format!(
373 "GET /hmr HTTP/1.1\r\n\
374 Host: demo.localhost:{proxy_port}\r\n\
375 Connection: keep-alive, UpGrAdE\r\n\
376 Upgrade: WebSocket\r\n\
377 Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
378 Sec-WebSocket-Version: 13\r\n\
379 \r\n"
380 )
381 .as_bytes(),
382 )
383 .await
384 .unwrap();
385 let response = String::from_utf8(read_until_headers(&mut client).await).unwrap();
386 assert!(response.starts_with("HTTP/1.1 101 Switching Protocols\r\n"));
387
388 client.write_all(b"ping").await.unwrap();
389 let mut pong = [0_u8; 4];
390 client.read_exact(&mut pong).await.unwrap();
391 assert_eq!(&pong, b"pong");
392
393 upstream.await.unwrap();
394 proxy.abort();
395 }
396}