use std::net::SocketAddr;
use std::time::Duration;
use bytes::Bytes;
use h2ts_server::{
accept_with_options, bridge_with, AcceptOptions, BridgeConfig, CloseFrame, KeepAlive,
};
use http_body_util::Empty;
use hyper::body::Incoming;
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Request, Response};
use hyper_util::rt::TokioIo;
use tokio::net::{TcpListener, TcpStream};
type BoxError = Box<dyn std::error::Error + Send + Sync>;
#[tokio::main]
async fn main() -> Result<(), BoxError> {
let raw: Vec<String> = std::env::args().skip(1).collect();
let allow_implicit_codec = raw.iter().any(|a| a == "--allow-implicit-codec");
let no_keepalive = raw.iter().any(|a| a == "--no-keepalive");
let mut allowed_origins: Vec<String> = Vec::new();
let mut positional: Vec<String> = Vec::new();
let mut it = raw.into_iter();
while let Some(a) = it.next() {
match a.as_str() {
"--allow-implicit-codec" | "--no-keepalive" => {}
"--allowed-origin" => {
if let Some(v) = it.next() {
allowed_origins.push(v);
}
}
_ => positional.push(a),
}
}
let allowed_origins = (!allowed_origins.is_empty()).then_some(allowed_origins);
let mut args = positional.into_iter();
let listen: SocketAddr = args
.next()
.unwrap_or_else(|| "127.0.0.1:8091".to_string())
.parse()?;
let upstream: String = args.next().unwrap_or_else(|| "127.0.0.1:8000".to_string());
let keepalive: Option<KeepAlive> = if no_keepalive {
None
} else {
match args.next() {
Some(s) => s
.parse::<u64>()
.ok()
.filter(|&s| s > 0)
.map(|s| KeepAlive::new(Duration::from_secs(s), Duration::from_secs(s))),
None => Some(KeepAlive::default()),
}
};
let listener = TcpListener::bind(listen).await?;
let ka = match &keepalive {
Some(k) => format!("keepalive {}s", k.interval.as_secs()),
None => "keepalive off".to_string(),
};
let codec = if allow_implicit_codec {
"any subprotocol"
} else {
"h2ts only"
};
let origins = match &allowed_origins {
Some(o) => format!("origins [{}]", o.join(", ")),
None => "any origin".to_string(),
};
eprintln!(
"[h2ts-proxy] listening ws://{listen} -> tcp://{upstream} (h2c, {ka}, {codec}, {origins})"
);
loop {
let (socket, peer) = match listener.accept().await {
Ok(pair) => pair,
Err(err) => {
eprintln!("[h2ts-proxy] accept error: {err}");
tokio::time::sleep(Duration::from_millis(50)).await;
continue;
}
};
let upstream = upstream.clone();
let keepalive = keepalive.clone();
let allowed_origins = allowed_origins.clone();
tokio::spawn(async move {
let io = TokioIo::new(socket);
let service = service_fn(move |req| {
handle(
req,
upstream.clone(),
peer,
keepalive.clone(),
allow_implicit_codec,
allowed_origins.clone(),
)
});
if let Err(err) = http1::Builder::new()
.serve_connection(io, service)
.with_upgrades()
.await
{
eprintln!("[h2ts-proxy] connection error ({peer}): {err}");
}
});
}
}
async fn handle(
mut req: Request<Incoming>,
upstream: String,
peer: SocketAddr,
keepalive: Option<KeepAlive>,
allow_implicit_codec: bool,
allowed_origins: Option<Vec<String>>,
) -> Result<Response<Empty<Bytes>>, BoxError> {
let (response, ws_fut) = match accept_with_options(
&mut req,
|_offered| None,
AcceptOptions {
allow_implicit_codec,
allowed_origins,
},
) {
Ok(pair) => pair,
Err(err) => {
eprintln!("[h2ts-proxy] rejected ({peer}): {err}");
return Ok(err.rejection_response());
}
};
tokio::spawn(async move {
let ws = match ws_fut.await {
Ok(ws) => ws,
Err(err) => {
eprintln!("[h2ts-proxy] ws upgrade failed ({peer}): {err}");
return;
}
};
let upstream_tcp = match TcpStream::connect(&upstream).await {
Ok(tcp) => tcp,
Err(err) => {
eprintln!("[h2ts-proxy] upstream connect failed ({upstream}): {err}");
let cfg = BridgeConfig {
close: CloseFrame {
code: 1014,
reason: "bad gateway".to_string(),
},
keepalive: None,
..Default::default()
};
let _ = bridge_with(ws, tokio::io::empty(), cfg).await;
return;
}
};
eprintln!("[h2ts-proxy] bridging ({peer}) <-> {upstream}");
let config = BridgeConfig {
keepalive,
error_close: CloseFrame {
code: 1014,
reason: "bad gateway".to_string(),
},
on_close: Some(Box::new(move |cf| {
eprintln!("[h2ts-proxy] closed ({peer}): {} {:?}", cf.code, cf.reason);
})),
..Default::default()
};
if let Err(err) = bridge_with(ws, upstream_tcp, config).await {
eprintln!("[h2ts-proxy] bridge error ({peer}): {err}");
}
});
Ok(response)
}