use std::sync::{
Arc,
atomic::{AtomicBool, Ordering},
};
use tokio::{
net::{TcpListener, TcpStream},
select, spawn,
sync::broadcast,
};
pub struct TcpProxy {
addr: String,
paused: Arc<AtomicBool>,
kill_tx: broadcast::Sender<()>,
}
impl TcpProxy {
pub async fn start(upstream_port: u16) -> Self {
let listener = TcpListener::bind("[::1]:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let paused = Arc::new(AtomicBool::new(false));
let (kill_tx, _) = broadcast::channel(16);
let accept_paused = paused.clone();
let accept_kill = kill_tx.clone();
spawn(async move {
loop {
let inbound = match listener.accept().await {
Ok((stream, _)) => stream,
Err(_) => continue,
};
if accept_paused.load(Ordering::SeqCst) {
drop(inbound);
continue;
}
spawn(relay(inbound, upstream_port, accept_kill.subscribe()));
}
});
Self {
addr: format!("[::1]:{}", port),
paused,
kill_tx,
}
}
pub fn addr(&self) -> &str {
&self.addr
}
pub fn ws_url(&self) -> String {
format!("ws://{}", self.addr)
}
pub fn kill(&self) {
let _ = self.kill_tx.send(());
}
pub fn pause(&self) {
self.paused.store(true, Ordering::SeqCst);
}
}
async fn relay(inbound: TcpStream, upstream_port: u16, mut kill_rx: broadcast::Receiver<()>) {
let upstream = match TcpStream::connect(format!("[::1]:{}", upstream_port)).await {
Ok(stream) => stream,
Err(_) => return,
};
let (mut client_read, mut client_write) = inbound.into_split();
let (mut server_read, mut server_write) = upstream.into_split();
select! {
_ = tokio::io::copy(&mut client_read, &mut server_write) => {}
_ = tokio::io::copy(&mut server_read, &mut client_write) => {}
_ = kill_rx.recv() => {}
}
}