1use std::collections::HashMap;2use std::net::SocketAddr;3use std::sync::atomic::{AtomicU64, Ordering};4use std::sync::{Arc, Mutex};56use bifrostlink::declarative::endpoints;7use bifrostlink::Config;8use remowt_link_shared::iroh_tunnel::{TunnelAddr, TunnelDialer};9use serde::{Deserialize, Serialize};10use std::result::Result;11use tokio::net::{TcpStream, UdpSocket};12use tracing::warn;1314#[derive(Serialize, Deserialize, Debug, thiserror::Error)]15pub enum Error {16 #[error("tunnel unavailable: {0}")]17 Tunnel(String),18 #[error("connect to {0} failed: {1}")]19 Connect(String, String),20 #[error("invalid target address {0:?}")]21 BadAddr(String),22 #[error("udp forward requires the iroh fast path, which is not established")]23 NoIroh,24}2526#[derive(Clone)]27pub struct Forward {28 dialer: Arc<TunnelDialer>,29 next_session: Arc<AtomicU64>,30}3132impl Forward {33 pub fn new(dialer: Arc<TunnelDialer>) -> Self {34 Self {35 dialer,36 next_session: Arc::new(AtomicU64::new(0)),37 }38 }39}4041#[endpoints(ns = 12)]42impl Forward {43 #[endpoints(id = 1)]44 async fn connect_tcp(&self, tunnel: TunnelAddr, addr: String) -> Result<(), Error> {45 let stream = self46 .dialer47 .connect_tunnel(&tunnel)48 .await49 .map_err(|e| Error::Tunnel(e.to_string()))?;50 let tcp = TcpStream::connect(&addr)51 .await52 .map_err(|e| Error::Connect(addr, e.to_string()))?;53 tokio::spawn(async move {54 let mut stream = stream;55 let mut tcp = tcp;56 let _ = tokio::io::copy_bidirectional(&mut stream, &mut tcp).await;57 });58 Ok(())59 }6061 #[endpoints(id = 2)]62 async fn open_udp(&self, addr: String) -> Result<u64, Error> {63 let target: SocketAddr = addr.parse().map_err(|_| Error::BadAddr(addr.clone()))?;64 let router = self.dialer.router().ok_or(Error::NoIroh)?;65 let session = self.next_session.fetch_add(1, Ordering::Relaxed);66 let mut rx = router.register(session);6768 let sockets: Arc<Mutex<HashMap<u64, Arc<UdpSocket>>>> =69 Arc::new(Mutex::new(HashMap::new()));70 tokio::spawn(async move {71 while let Some((sub, payload)) = rx.recv().await {72 let existing = sockets.lock().expect("lock").get(&sub).cloned();73 let sock = match existing {74 Some(s) => s,75 None => {76 let sock = match UdpSocket::bind(unspecified_for(&target)).await {77 Ok(s) => s,78 Err(e) => {79 warn!("udp forward: bind failed: {e}");80 continue;81 }82 };83 if let Err(e) = sock.connect(target).await {84 warn!("udp forward: connect {target} failed: {e}");85 continue;86 }87 let sock = Arc::new(sock);88 sockets.lock().expect("lock").insert(sub, sock.clone());89 90 let router = router.clone();91 let reply_sock = sock.clone();92 tokio::spawn(async move {93 let mut buf = vec![0u8; 65535];94 while let Ok(n) = reply_sock.recv(&mut buf).await {95 if router.send(session, sub, &buf[..n]).is_err() {96 break;97 }98 }99 });100 sock101 }102 };103 let _ = sock.send(&payload).await;104 }105 router.unregister(session);106 });107108 Ok(session)109 }110}111112fn unspecified_for(target: &SocketAddr) -> SocketAddr {113 match target {114 SocketAddr::V4(_) => SocketAddr::from(([0, 0, 0, 0], 0)),115 SocketAddr::V6(_) => SocketAddr::from(([0u16; 8], 0)),116 }117}