Skip to main content

stackless_daemon/
proxy.rs

1//! The built-in reverse proxy (§3): HTTP plus transparent HTTP/1
2//! upgrade tunnels, one fixed unprivileged port, routing on the Host
3//! header. Listens on both loopbacks — macOS resolves multi-label
4//! `*.localhost` to `::1` (verified 2026-06-11), while the health
5//! checker dials 127.0.0.1 with an explicit Host header.
6
7use 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
24/// The proxy's default port (D9) — configurable globally via
25/// `STACKLESS_PROXY_PORT`, never per instance: origins must stay
26/// derivable from the instance name alone.
27pub 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/// Either an upstream body streamed through, or a local message.
177#[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        // SAFETY-free projection: match on the enum through get_mut.
192        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}