use std::io;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio::net::{TcpListener, TcpStream};
use boatramp_core::deploy::DeployStore;
use boatramp_core::gateway::Upstream;
use boatramp_core::route::{self, Outcome};
use boatramp_core::security::SecurityPosture;
use crate::proxy;
const MAX_HEAD: usize = 16 * 1024;
const HOP_BY_HOP: &[&str] = &[
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
];
#[derive(Clone)]
pub struct SpliceCtx {
pub deploy: DeployStore,
pub posture: SecurityPosture,
pub daemon: Option<Arc<crate::DaemonRuntime>>,
}
struct SplicePlan {
resolved: proxy::ResolvedTarget,
upstream: Upstream,
site: String,
project: String,
}
const PEEK_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
pub async fn serve<S>(
tcp: TcpListener,
ctx: SpliceCtx,
router: axum::Router,
shutdown: S,
) -> io::Result<()>
where
S: std::future::Future<Output = ()> + Send,
{
tokio::pin!(shutdown);
loop {
tokio::select! {
_ = &mut shutdown => return Ok(()),
accepted = tcp.accept() => {
let (mut io, peer) = match accepted {
Ok(v) => v,
Err(err) => {
tracing::debug!(%err, "splice serve: accept error");
continue;
}
};
crate::disable_nagle(&mut io);
let ctx = ctx.clone();
let router = router.clone();
tokio::spawn(async move {
let eligible = tokio::time::timeout(PEEK_TIMEOUT, peek_classify(&io, &ctx))
.await
.unwrap_or_default();
match eligible {
Some(plan) => {
if let Err(err) = splice_conn(io, peer, plan, ctx, router).await {
tracing::debug!(%peer, %err, "splice connection ended");
}
}
None => serve_fallback(io, peer, router).await,
}
});
}
}
}
}
struct Rewind {
pre: Vec<u8>,
pos: usize,
inner: TcpStream,
}
impl tokio::io::AsyncRead for Rewind {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<io::Result<()>> {
if self.pos < self.pre.len() {
let n = (self.pre.len() - self.pos).min(buf.remaining());
let start = self.pos;
buf.put_slice(&self.pre[start..start + n]);
self.pos += n;
return std::task::Poll::Ready(Ok(()));
}
std::pin::Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
impl tokio::io::AsyncWrite for Rewind {
fn poll_write(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &[u8],
) -> std::task::Poll<io::Result<usize>> {
std::pin::Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<io::Result<()>> {
std::pin::Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<io::Result<()>> {
std::pin::Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
async fn serve_fallback<IO>(io: IO, peer: SocketAddr, router: axum::Router)
where
IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let io = hyper_util::rt::TokioIo::new(io);
let svc = hyper::service::service_fn(move |mut req: hyper::Request<hyper::body::Incoming>| {
req.extensions_mut()
.insert(axum::extract::ConnectInfo(peer));
let mut router = router.clone();
async move {
use tower_service::Service as _;
router.call(req).await
}
});
if let Err(err) = hyper::server::conn::http1::Builder::new()
.serve_connection(io, svc)
.with_upgrades()
.await
{
tracing::debug!(%peer, %err, "connection served with error");
}
}
async fn fall_back(
head: Vec<u8>,
leftover: Vec<u8>,
client: TcpStream,
peer: SocketAddr,
router: axum::Router,
) -> io::Result<()> {
let mut pre = head;
pre.extend_from_slice(&leftover);
let rewind = Rewind {
pre,
pos: 0,
inner: client,
};
serve_fallback(rewind, peer, router).await;
Ok(())
}
async fn peek_classify(io: &TcpStream, ctx: &SpliceCtx) -> Option<SplicePlan> {
let mut buf = vec![0u8; MAX_HEAD];
loop {
io.readable().await.ok()?;
let n = match io.peek(&mut buf).await {
Ok(0) => return None, Ok(n) => n,
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => continue,
Err(_) => return None,
};
if let Some(end) = find_head_end(&buf[..n]) {
return classify(&buf[..end], ctx).await;
}
if n >= MAX_HEAD {
return None; }
tokio::task::yield_now().await;
}
}
async fn classify(head: &[u8], ctx: &SpliceCtx) -> Option<SplicePlan> {
if !splice_supported() {
return None;
}
let (method, target_path, headers) = parse_request_head(head)?;
if !matches!(method.as_str(), "GET" | "HEAD") {
return None;
}
if header(&headers, "content-length").is_some()
|| header(&headers, "transfer-encoding").is_some()
|| header(&headers, "upgrade").is_some()
|| header(&headers, "expect").is_some()
{
return None;
}
let host = header(&headers, "host").map(strip_port).unwrap_or("");
let path = path_only(&target_path);
if is_reserved_path(path) {
return None;
}
let eff = ctx.daemon.as_ref().map(|d| d.effective());
#[cfg(feature = "console")]
if let Some(eff) = &eff {
if crate::console::would_intercept(eff, host, path) {
return None;
}
}
let (project, site) = match ctx.deploy.resolve_site_by_host(host).await.ok()? {
Some(owner) => (owner.project, owner.site),
None => (
boatramp_core::project::ProjectRef::DEFAULT
.as_str()
.to_string(),
eff.as_ref().and_then(|e| e.default_site.clone())?,
),
};
let project_ref = boatramp_core::project::ProjectRef::new(&project);
let cfg = ctx
.deploy
.get_site_config_cached(project_ref, &site)
.await
.ok()??;
let a = &cfg.access;
let permissive = a.basic_auth.is_none()
&& a.rate_limit.is_none()
&& a.ip.allow.is_empty()
&& a.ip.deny.is_empty()
&& a.trusted_proxies.is_empty()
&& a.waf == Default::default();
if !permissive {
return None;
}
if boatramp_core::config::transport_redirect(
&cfg.security,
&cfg.domains,
"http",
host,
&target_path,
)
.is_some()
{
return None;
}
if cfg.handlers.is_some() {
return None;
}
let manifest = ctx
.deploy
.current_manifest(project_ref, &site)
.await
.ok()??;
let ctx_default = boatramp_core::predicate::RequestContext::default();
let resolved = route::resolve_ctx(&manifest.config, &manifest.files, path, &ctx_default);
match resolved.outcome {
Outcome::Redirect { .. } | Outcome::Proxy { .. } => return None,
Outcome::File { .. } | Outcome::NotFound { .. } => {}
}
if route::match_handler(&manifest.config.handlers, &method, path).is_some() {
return None;
}
if !manifest.config.streams.is_empty() {
return None;
}
let gw = cfg.gateway.as_ref().filter(|g| g.is_enabled())?;
let route_match = gw.match_route(path)?;
let upstream = gw.upstreams.get(&route_match.upstream)?;
if upstream.compute.is_some() || upstream.discover.is_some() || upstream.active_health.is_some()
{
return None;
}
let backends = upstream.static_backends();
if backends.len() != 1 {
return None;
}
let target = backends[0];
if !target.starts_with("http://") {
return None;
}
let resolved_target = proxy::resolve_target(target, &ctx.posture).await.ok()?;
if !proxy::gateway_addr_allowed(resolved_target.addr.ip(), &ctx.posture) {
return None;
}
Some(SplicePlan {
resolved: resolved_target,
upstream: upstream.clone(),
site,
project,
})
}
async fn splice_conn(
mut client: TcpStream,
peer: SocketAddr,
plan: SplicePlan,
ctx: SpliceCtx,
router: axum::Router,
) -> io::Result<()> {
let mut upstream = TcpStream::connect(plan.resolved.addr).await?;
upstream.set_nodelay(true).ok();
let client_ip: IpAddr = peer.ip();
loop {
let (head, leftover) = match read_head(&mut client).await? {
Some(v) => v,
None => return Ok(()), };
let eligible = parse_request_head(&head)
.filter(|(m, _, _)| matches!(m.as_str(), "GET" | "HEAD") && leftover.is_empty());
let (method, target_path, headers) = match eligible {
Some(v) => v,
None => return fall_back(head, leftover, client, peer, router).await,
};
let client_close = header(&headers, "connection")
.is_some_and(|v| v.to_ascii_lowercase().contains("close"));
let path = path_only(&target_path);
if plan
.upstream_route(&ctx, &plan.project, &plan.site, path)
.await
.is_none()
{
return fall_back(head, leftover, client, peer, router).await;
}
let up_head = build_upstream_head(&plan, &method, &target_path, &headers, client_ip);
upstream.write_all(&up_head).await?;
let (resp_head, body_prefix) = match read_head(&mut upstream).await? {
Some(v) => v,
None => return Ok(()),
};
let (status_ok, content_length, chunked, close) = parse_response_head(&resp_head);
let head_only = method == "HEAD" || matches!(status_ok, Some(204) | Some(304));
let out_head = rewrite_response_head(&resp_head, &plan.upstream);
client.write_all(&out_head).await?;
if !body_prefix.is_empty() && !head_only {
client.write_all(&body_prefix).await?;
}
if head_only {
if close || client_close {
return Ok(());
}
continue;
}
match content_length {
Some(total) if !chunked => {
let remaining = total.saturating_sub(body_prefix.len());
if remaining > 0 {
splice_body(&upstream, &client, remaining).await?;
}
}
_ => {
relay_to_close(&mut upstream, &mut client).await?;
return Ok(());
}
}
if close || client_close {
return Ok(());
}
}
}
impl SplicePlan {
async fn upstream_route(
&self,
ctx: &SpliceCtx,
project: &str,
site: &str,
path: &str,
) -> Option<()> {
let project_ref = boatramp_core::project::ProjectRef::new(project);
let cfg = ctx
.deploy
.get_site_config_cached(project_ref, site)
.await
.ok()??;
let gw = cfg.gateway.as_ref().filter(|g| g.is_enabled())?;
let route_match = gw.match_route(path)?;
let upstream = gw.upstreams.get(&route_match.upstream)?;
let backends = upstream.static_backends();
if backends.len() == 1 && backends[0].starts_with("http://") {
Some(())
} else {
None
}
}
}
type RequestHead = (String, String, Vec<(String, String)>);
fn parse_request_head(head: &[u8]) -> Option<RequestHead> {
let text = std::str::from_utf8(head).ok()?;
let mut lines = text.split("\r\n");
let request_line = lines.next()?;
let mut parts = request_line.split(' ');
let method = parts.next()?.to_string();
let target = parts.next()?.to_string();
let version = parts.next()?;
if version != "HTTP/1.1" {
return None; }
let mut headers = Vec::new();
for line in lines {
if line.is_empty() {
break;
}
if let Some((k, v)) = line.split_once(':') {
headers.push((k.trim().to_string(), v.trim().to_string()));
}
}
Some((method, target, headers))
}
fn parse_response_head(head: &[u8]) -> (Option<u16>, Option<usize>, bool, bool) {
let text = match std::str::from_utf8(head) {
Ok(t) => t,
Err(_) => return (None, None, false, true),
};
let mut lines = text.split("\r\n");
let status_line = lines.next().unwrap_or("");
let status: Option<u16> = status_line.split(' ').nth(1).and_then(|s| s.parse().ok());
let mut content_length = None;
let mut chunked = false;
let mut close = false;
for line in lines {
if line.is_empty() {
break;
}
if let Some((k, v)) = line.split_once(':') {
let k = k.trim();
if k.eq_ignore_ascii_case("content-length") {
content_length = v.trim().parse().ok();
} else if k.eq_ignore_ascii_case("transfer-encoding")
&& v.to_ascii_lowercase().contains("chunked")
{
chunked = true;
} else if k.eq_ignore_ascii_case("connection")
&& v.to_ascii_lowercase().contains("close")
{
close = true;
}
}
}
(status, content_length, chunked, close)
}
fn build_upstream_head(
plan: &SplicePlan,
method: &str,
target_path: &str,
headers: &[(String, String)],
client_ip: IpAddr,
) -> Vec<u8> {
let base = plan.resolved.parsed.path().trim_end_matches('/');
let (req_path, query) = match target_path.split_once('?') {
Some((p, q)) => (p, Some(q)),
None => (target_path, None),
};
let forwarded = plan.upstream.forward_path(req_path);
let mut out = format!("{method} {base}{forwarded}");
if let Some(q) = query {
out.push('?');
out.push_str(q);
}
out.push_str(" HTTP/1.1\r\n");
let host = plan
.upstream
.host_header
.as_deref()
.unwrap_or(plan.resolved.host.as_str());
out.push_str(&format!("host: {host}\r\n"));
for (name, value) in headers {
let lname = name.to_ascii_lowercase();
if lname == "host"
|| HOP_BY_HOP.contains(&lname.as_str())
|| plan
.upstream
.header_up
.remove
.iter()
.any(|h| lname.eq_ignore_ascii_case(h))
{
continue;
}
out.push_str(&format!("{name}: {value}\r\n"));
}
out.push_str(&format!("x-forwarded-for: {client_ip}\r\n"));
out.push_str("x-forwarded-proto: http\r\n");
if let Some(h) = headers.iter().find(|(k, _)| k.eq_ignore_ascii_case("host")) {
out.push_str(&format!("x-forwarded-host: {}\r\n", h.1));
}
for (name, value) in &plan.upstream.header_up.set {
out.push_str(&format!("{name}: {value}\r\n"));
}
out.push_str("\r\n");
out.into_bytes()
}
fn rewrite_response_head(resp_head: &[u8], upstream: &Upstream) -> Vec<u8> {
let text = match std::str::from_utf8(resp_head) {
Ok(t) => t,
Err(_) => return resp_head.to_vec(),
};
let mut lines = text.split("\r\n");
let status_line = lines.next().unwrap_or("HTTP/1.1 502 Bad Gateway");
let mut out = String::with_capacity(resp_head.len());
out.push_str(status_line);
out.push_str("\r\n");
for line in lines {
if line.is_empty() {
break;
}
if let Some((k, _)) = line.split_once(':') {
let lk = k.trim().to_ascii_lowercase();
if HOP_BY_HOP.contains(&lk.as_str())
|| upstream
.header_down
.remove
.iter()
.any(|h| lk.eq_ignore_ascii_case(h))
{
continue;
}
}
out.push_str(line);
out.push_str("\r\n");
}
for (name, value) in &upstream.header_down.set {
out.push_str(&format!("{name}: {value}\r\n"));
}
out.push_str("\r\n");
out.into_bytes()
}
async fn read_head(s: &mut TcpStream) -> io::Result<Option<(Vec<u8>, Vec<u8>)>> {
use tokio::io::AsyncReadExt;
let mut acc: Vec<u8> = Vec::with_capacity(1024);
let mut tmp = [0u8; 8192];
loop {
let n = s.read(&mut tmp).await?;
if n == 0 {
return if acc.is_empty() {
Ok(None)
} else {
Err(io::Error::new(io::ErrorKind::UnexpectedEof, "eof mid-head"))
};
}
acc.extend_from_slice(&tmp[..n]);
if let Some(end) = find_head_end(&acc) {
let head = acc[..end].to_vec();
let rest = acc[end..].to_vec();
return Ok(Some((head, rest)));
}
if acc.len() > MAX_HEAD {
return Err(io::Error::new(io::ErrorKind::InvalidData, "head too large"));
}
}
}
fn find_head_end(b: &[u8]) -> Option<usize> {
b.windows(4).position(|w| w == b"\r\n\r\n").map(|p| p + 4)
}
async fn relay_to_close(src: &mut TcpStream, dst: &mut TcpStream) -> io::Result<()> {
tokio::io::copy(src, dst).await.map(|_| ())
}
#[cfg(target_os = "linux")]
async fn splice_body(src: &TcpStream, dst: &TcpStream, mut n: usize) -> io::Result<()> {
use std::os::fd::AsRawFd;
use std::ptr;
use tokio::io::Interest;
let mut fds = [0i32; 2];
if unsafe { libc::pipe2(fds.as_mut_ptr(), libc::O_NONBLOCK) } != 0 {
return Err(io::Error::last_os_error());
}
let (pr, pw) = (fds[0], fds[1]);
let src_fd = src.as_raw_fd();
let dst_fd = dst.as_raw_fd();
let flags = libc::SPLICE_F_MOVE | libc::SPLICE_F_NONBLOCK;
let result = async {
while n > 0 {
let want = n.min(1 << 16);
let in_n = src
.async_io(Interest::READABLE, || {
let r = unsafe {
libc::splice(src_fd, ptr::null_mut(), pw, ptr::null_mut(), want, flags)
};
if r < 0 {
Err(io::Error::last_os_error())
} else {
Ok(r as usize)
}
})
.await?;
if in_n == 0 {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"upstream closed before content-length",
));
}
let mut left = in_n;
while left > 0 {
let out_n = dst
.async_io(Interest::WRITABLE, || {
let r = unsafe {
libc::splice(pr, ptr::null_mut(), dst_fd, ptr::null_mut(), left, flags)
};
if r < 0 {
Err(io::Error::last_os_error())
} else {
Ok(r as usize)
}
})
.await?;
left -= out_n;
}
n -= in_n;
}
Ok(())
}
.await;
unsafe {
libc::close(pr);
libc::close(pw);
}
result
}
#[cfg(not(target_os = "linux"))]
async fn splice_body(_src: &TcpStream, _dst: &TcpStream, _n: usize) -> io::Result<()> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"splice is linux-only",
))
}
pub fn splice_supported() -> bool {
cfg!(target_os = "linux")
}
fn header<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> {
headers
.iter()
.find(|(k, _)| k.eq_ignore_ascii_case(name))
.map(|(_, v)| v.as_str())
}
fn strip_port(host: &str) -> &str {
host.rsplit_once(':')
.filter(|(h, _)| !h.contains(':') || h.starts_with('['))
.map(|(h, _)| h)
.unwrap_or(host)
.trim_start_matches('[')
.trim_end_matches(']')
}
fn path_only(target: &str) -> &str {
target.split('?').next().unwrap_or(target)
}
fn is_reserved_path(path: &str) -> bool {
const RESERVED: &[&str] = &[
"/healthz",
"/readyz",
"/api/",
"/_sites/",
"/_deploy/",
"/_webhooks/",
"/.well-known/boatramp-",
"/mcp",
];
RESERVED.iter().any(|p| path.starts_with(p))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reserved_paths_defer_to_the_router() {
for p in [
"/api/sites",
"/healthz",
"/readyz",
"/_sites/x",
"/_deploy/abc",
"/_webhooks/hook",
"/.well-known/boatramp-domain-verification/tok",
"/mcp",
] {
assert!(is_reserved_path(p), "{p} should be reserved");
}
for p in [
"/",
"/b/100k",
"/index.html",
"/assets/app.js",
"/.well-known/acme-challenge/x",
] {
assert!(!is_reserved_path(p), "{p} should not be reserved");
}
}
#[test]
fn request_head_parse_only_http11_bodyless() {
let (m, t, h) = parse_request_head(
b"GET /b/100k?x=1 HTTP/1.1\r\nHost: a.example\r\nAccept: */*\r\n\r\n",
)
.expect("valid GET");
assert_eq!(m, "GET");
assert_eq!(t, "/b/100k?x=1");
assert_eq!(header(&h, "host"), Some("a.example"));
assert!(parse_request_head(b"GET / HTTP/1.0\r\nHost: a\r\n\r\n").is_none());
}
#[test]
fn response_head_parse_detects_length_chunked_close() {
let (st, content_length, chunked, close) = parse_response_head(
b"HTTP/1.1 200 OK\r\nContent-Length: 102400\r\nContent-Type: x\r\n\r\n",
);
assert_eq!(st, Some(200));
assert_eq!(content_length, Some(102400));
assert!(!chunked && !close);
let (_, content_length2, chunked2, close2) = parse_response_head(
b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n",
);
assert_eq!(content_length2, None);
assert!(chunked2 && close2);
}
#[test]
fn strip_port_handles_ipv6_and_default() {
assert_eq!(strip_port("host.example:8080"), "host.example");
assert_eq!(strip_port("host.example"), "host.example");
assert_eq!(strip_port("[::1]:80"), "::1");
assert_eq!(strip_port("127.0.0.1:9000"), "127.0.0.1");
}
#[test]
fn response_head_rewrite_drops_hop_by_hop_keeps_length() {
let up = Upstream::default();
let out = rewrite_response_head(
b"HTTP/1.1 200 OK\r\nContent-Length: 3\r\nConnection: keep-alive\r\nTransfer-Encoding: chunked\r\nContent-Type: text/plain\r\n\r\n",
&up,
);
let text = String::from_utf8(out).unwrap();
assert!(text.starts_with("HTTP/1.1 200 OK\r\n"));
assert!(text.contains("content-length: 3\r\n") || text.contains("Content-Length: 3\r\n"));
assert!(text
.to_ascii_lowercase()
.contains("content-type: text/plain"));
assert!(!text.to_ascii_lowercase().contains("connection:"));
assert!(!text.to_ascii_lowercase().contains("transfer-encoding:"));
}
#[test]
fn head_end_index() {
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n\r\nBODY"), Some(18));
assert_eq!(find_head_end(b"partial\r\nno end"), None);
}
#[tokio::test]
async fn serve_loop_proxies_gateway_and_defers_reserved_routes() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let up = TcpListener::bind("127.0.0.1:0").await.unwrap();
let up_addr = up.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut s, _)) = up.accept().await {
tokio::spawn(async move {
let mut buf = [0u8; 4096];
loop {
match s.read(&mut buf).await {
Ok(0) | Err(_) => return,
Ok(n) if !buf[..n].windows(4).any(|w| w == b"\r\n\r\n") => continue,
Ok(_) => {}
}
let body = b"hello-from-upstream";
let head =
format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
if s.write_all(head.as_bytes()).await.is_err()
|| s.write_all(body).await.is_err()
{
return;
}
}
});
}
});
let addr = spawn_gateway_serve(up_addr).await;
async fn req(addr: SocketAddr, raw: &str) -> String {
let mut c = TcpStream::connect(addr).await.unwrap();
c.write_all(raw.as_bytes()).await.unwrap();
let mut out = Vec::new();
c.read_to_end(&mut out).await.unwrap();
String::from_utf8_lossy(&out).into_owned()
}
let resp = req(
addr,
"GET /anything HTTP/1.1\r\nHost: test.local\r\nConnection: close\r\n\r\n",
)
.await;
assert!(
resp.contains("hello-from-upstream"),
"proxy body missing: {resp}"
);
let hz = req(
addr,
"GET /healthz HTTP/1.1\r\nHost: test.local\r\nConnection: close\r\n\r\n",
)
.await;
assert!(
hz.starts_with("HTTP/1.1 200") && hz.to_ascii_lowercase().contains("ok"),
"healthz fallback: {hz}"
);
}
async fn spawn_gateway_serve(up_addr: SocketAddr) -> SocketAddr {
use boatramp_core::config::{DomainConfig, SiteConfig};
use boatramp_core::deploy::{DeployStore, Manifest};
use boatramp_core::gateway::{GatewayConfig, GatewayRoute, Upstream};
use boatramp_core::kv::MemoryKv;
use boatramp_core::project::ProjectRef;
use boatramp_core::security::SecurityProfile;
let deploy = DeployStore::new(
Arc::new(boatramp_storage::FsStorage::new(std::env::temp_dir())),
Arc::new(MemoryKv::new()),
);
let cfg = SiteConfig {
domains: DomainConfig {
primary: Some("test.local".into()),
..Default::default()
},
gateway: Some(GatewayConfig {
upstreams: std::iter::once((
"backend".to_string(),
Upstream {
target: format!("http://{up_addr}"),
..Default::default()
},
))
.collect(),
routes: vec![GatewayRoute {
matches: "/**".into(),
upstream: "backend".into(),
}],
}),
..Default::default()
};
deploy
.set_site_config(ProjectRef::DEFAULT, "www", &cfg)
.await
.unwrap();
let id = deploy.put_manifest(&Manifest::default()).await.unwrap();
deploy
.activate(ProjectRef::DEFAULT, "www", &id)
.await
.unwrap();
let posture = SecurityProfile::Dev.preset();
let router = crate::router_with(
deploy.clone(),
crate::Auth::disabled(),
crate::HandlerRuntime::disabled(),
crate::ServerOptions {
posture,
..Default::default()
},
);
let ctx = SpliceCtx {
deploy,
posture,
daemon: None,
};
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let _ = serve(listener, ctx, router, std::future::pending::<()>()).await;
});
addr
}
async fn read_one_response(c: &mut TcpStream) -> String {
use tokio::io::AsyncReadExt;
let mut buf = Vec::new();
let mut tmp = [0u8; 1024];
let head_end = loop {
if let Some(p) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break p + 4;
}
let n = c.read(&mut tmp).await.unwrap();
if n == 0 {
return String::from_utf8_lossy(&buf).into_owned();
}
buf.extend_from_slice(&tmp[..n]);
};
let content_length = String::from_utf8_lossy(&buf[..head_end])
.to_ascii_lowercase()
.split("\r\n")
.find_map(|l| l.strip_prefix("content-length:").map(str::to_owned))
.and_then(|v| v.trim().parse::<usize>().ok())
.unwrap_or(0);
while buf.len() < head_end + content_length {
let n = c.read(&mut tmp).await.unwrap();
if n == 0 {
break;
}
buf.extend_from_slice(&tmp[..n]);
}
String::from_utf8_lossy(&buf).into_owned()
}
#[tokio::test]
async fn upstream_dying_mid_body_closes_the_client_instead_of_hanging() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::time::{timeout, Duration};
let up = TcpListener::bind("127.0.0.1:0").await.unwrap();
let up_addr = up.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut s, _)) = up.accept().await {
tokio::spawn(async move {
let mut buf = [0u8; 4096];
loop {
match s.read(&mut buf).await {
Ok(0) | Err(_) => return,
Ok(n) if !buf[..n].windows(4).any(|w| w == b"\r\n\r\n") => continue,
Ok(_) => break,
}
}
let _ = s
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 1048576\r\n\r\ntruncated-body!!",
)
.await;
});
}
});
let addr = spawn_gateway_serve(up_addr).await;
let mut c = TcpStream::connect(addr).await.unwrap();
c.write_all(b"GET /anything HTTP/1.1\r\nHost: test.local\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut out = Vec::new();
let read = timeout(Duration::from_secs(5), c.read_to_end(&mut out)).await;
assert!(
read.is_ok(),
"client was left hanging after the upstream truncated the body"
);
#[cfg(target_os = "linux")]
{
let text = String::from_utf8_lossy(&out);
assert!(
text.contains("truncated-body!!"),
"splice path must relay the partial body before closing: {text}"
);
}
}
#[tokio::test]
async fn non_eligible_request_after_spliced_get_falls_back_not_dropped() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::time::{timeout, Duration};
let up = TcpListener::bind("127.0.0.1:0").await.unwrap();
let up_addr = up.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((mut s, _)) = up.accept().await {
tokio::spawn(async move {
let mut buf = [0u8; 4096];
loop {
match s.read(&mut buf).await {
Ok(0) | Err(_) => return,
Ok(n) if !buf[..n].windows(4).any(|w| w == b"\r\n\r\n") => continue,
Ok(_) => {}
}
let body = b"ok";
let head =
format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
if s.write_all(head.as_bytes()).await.is_err()
|| s.write_all(body).await.is_err()
{
return;
}
}
});
}
});
let addr = spawn_gateway_serve(up_addr).await;
let mut c = TcpStream::connect(addr).await.unwrap();
c.write_all(b"GET /anything HTTP/1.1\r\nHost: test.local\r\n\r\n")
.await
.unwrap();
let first = read_one_response(&mut c).await;
assert!(
first.starts_with("HTTP/1.1 200"),
"GET not proxied: {first}"
);
c.write_all(b"POST /anything HTTP/1.1\r\nHost: test.local\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut rest = Vec::new();
let read = timeout(Duration::from_secs(5), c.read_to_end(&mut rest)).await;
assert!(
read.is_ok(),
"POST after a spliced GET was dropped (client hung)"
);
let text = String::from_utf8_lossy(&rest);
assert!(
text.starts_with("HTTP/1.1 "),
"POST after a spliced GET got no HTTP response: {text}"
);
}
proptest::proptest! {
#[test]
fn parsers_never_panic_on_arbitrary_bytes(data: Vec<u8>) {
let _ = parse_request_head(&data);
let _ = parse_response_head(&data);
let up = Upstream::default();
let _ = rewrite_response_head(&data, &up);
let _ = find_head_end(&data);
}
#[test]
fn parsers_never_panic_on_arbitrary_ascii(s in ".{0,4096}") {
let _ = parse_request_head(s.as_bytes());
let _ = parse_response_head(s.as_bytes());
}
#[test]
fn cl_plus_te_response_is_never_treated_as_pure_length(len in 0usize..1_000_000) {
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {len}\r\nTransfer-Encoding: chunked\r\n\r\n"
);
let (_, _content_length, chunked, _close) = parse_response_head(head.as_bytes());
proptest::prop_assert!(chunked, "CL+TE must be flagged chunked (desync guard)");
}
#[test]
fn wellformed_get_head_roundtrips(path in "/[a-zA-Z0-9_/.-]{0,64}", host in "[a-z][a-z0-9.-]{0,32}") {
let head = format!("GET {path} HTTP/1.1\r\nHost: {host}\r\nAccept: */*\r\n\r\n");
let parsed = parse_request_head(head.as_bytes());
proptest::prop_assert!(parsed.is_some());
let (m, t, h) = parsed.unwrap();
proptest::prop_assert_eq!(m, "GET");
proptest::prop_assert_eq!(t, path);
proptest::prop_assert_eq!(header(&h, "host"), Some(host.as_str()));
}
}
}