use std::net::SocketAddr;
use std::time::Duration;
use bytes::Bytes;
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};
use h2ts_server::{accept_with_options, bridge_with, AcceptOptions, BridgeConfig, KeepAlive};
type BoxError = Box<dyn std::error::Error + Send + Sync>;
#[tokio::main]
async fn main() -> Result<(), BoxError> {
let mut args: Vec<String> = std::env::args().skip(1).collect();
let allow_implicit_codec = args.iter().any(|a| a == "--allow-implicit-codec");
args.retain(|a| a != "--allow-implicit-codec");
let mut args = args.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> = args
.next()
.and_then(|s| s.parse::<u64>().ok())
.filter(|&s| s > 0)
.map(|s| KeepAlive::new(Duration::from_secs(s), Duration::from_secs(s)));
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" };
eprintln!("[h2ts-proxy] listening ws://{listen} -> tcp://{upstream} (h2c, {ka}, {codec})");
loop {
let (socket, peer) = listener.accept().await?;
let upstream = upstream.clone();
let keepalive = keepalive.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)
});
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,
) -> Result<Response<Empty<Bytes>>, BoxError> {
let (response, ws_fut) = match accept_with_options(
&mut req,
|_offered| None,
AcceptOptions {
allow_implicit_codec,
},
) {
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}");
return;
}
};
eprintln!("[h2ts-proxy] bridging ({peer}) <-> {upstream}");
let config = BridgeConfig {
keepalive,
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)
}