use std::{collections::HashMap, net::SocketAddr, sync::Arc, sync::OnceLock};
use anyhow::Context;
use futures_concurrency::future::TryJoin as _;
use sillad::{
listener::Listener,
tcp::{TcpListener, TcpPipe},
};
use tokio::sync::Mutex;
use crate::{bound_dialer::connect_addrs, litecopy::litecopy};
static CACHE: OnceLock<Mutex<HashMap<Vec<SocketAddr>, SocketAddr>>> = OnceLock::new();
pub async fn forward_addrs(mut dests: Vec<SocketAddr>) -> anyhow::Result<SocketAddr> {
dests.sort();
dests.dedup();
anyhow::ensure!(!dests.is_empty(), "no upstream addresses to forward to");
let cache = CACHE.get_or_init(|| Mutex::new(HashMap::new()));
let mut guard = cache.lock().await;
if let Some(addr) = guard.get(&dests) {
return Ok(*addr);
}
let listener = TcpListener::bind("127.0.0.1:0".parse().unwrap())
.await
.context("bind loopback forwarder")?;
let local = listener.local_addr().await;
spawn_forwarder(listener, dests.clone());
guard.insert(dests, local);
Ok(local)
}
fn spawn_forwarder(mut listener: TcpListener, dests: Vec<SocketAddr>) {
let dests = Arc::new(dests);
geph5_rt::spawn(async move {
loop {
let downstream = match listener.accept().await {
Ok(conn) => conn,
Err(err) => {
tracing::warn!(err = %err, "broker egress forwarder stopped accepting");
return;
}
};
let dests = dests.clone();
geph5_rt::spawn(async move {
if let Err(err) = splice(downstream, &dests).await {
tracing::debug!(err = %err, "broker egress forwarder connection ended");
}
})
.detach();
}
})
.detach();
}
async fn splice(downstream: TcpPipe, dests: &[SocketAddr]) -> anyhow::Result<()> {
let upstream = connect_addrs(dests)
.await
.context("forwarder upstream dial failed")?;
let (read_down, write_down) = tokio::io::split(downstream);
let (read_up, write_up) = tokio::io::split(upstream);
(litecopy(read_down, write_up), litecopy(read_up, write_down))
.try_join()
.await?;
Ok(())
}