borer-core 0.5.9

network borer
Documentation
use std::net::SocketAddr;

use anyhow::Context as _;
use tokio::io::{AsyncRead, AsyncWrite, copy_bidirectional};

use crate::{
    dial::{Dial as _, DirectDial},
    proto::{padding::Padding, trojan},
    stream::stats::StatsStream,
};

/// Handles a single inbound Trojan connection.
pub struct TrojanConnection<T> {
    inner: T,
    user: String,
    peer_addr: SocketAddr,
}

impl<T> TrojanConnection<T>
where
    T: AsyncRead + AsyncWrite + Unpin,
{
    /// Create a Trojan connection wrapper for the given authenticated user.
    pub fn new(ts: T, user: String, peer_addr: SocketAddr) -> Self {
        Self {
            inner: ts,
            user,
            peer_addr,
        }
    }

    /// Proxy the Trojan stream to its requested upstream address.
    pub async fn handle(mut self) -> anyhow::Result<()> {
        let stream = &mut self.inner;
        let req = trojan::Request::read_from(stream)
            .await
            .context("trojan request read failed")?;
        let addr = req.address.clone();

        let padding = req.is_padding();
        if padding {
            Padding::default().write_to(stream).await?;
        }

        debug!("trojan start connect {addr}");

        let mut stats_stream = StatsStream::new(
            stream,
            self.user,
            req.hash,
            self.peer_addr.to_string(),
            req.address.to_string(),
            padding,
        );
        let mut remote_ts = DirectDial::new(std::time::Duration::from_secs(3))
            .dial(addr.clone())
            .await
            .context(format!("failed to dial {addr}"))?;

        info!("[{padding}] Trojan Connect to {addr}. ",);
        if let Ok((a, b)) = copy_bidirectional(&mut stats_stream, &mut remote_ts).await {
            debug!(
                "trojan copy end for {} traffic: {}<=>{} total: {}",
                req.address,
                a,
                b,
                a + b
            );
        }
        Ok(())
    }
}

#[cfg(test)]
mod tests {
    use std::{
        net::{IpAddr, Ipv4Addr, SocketAddr},
        time::{SystemTime, UNIX_EPOCH},
    };

    use bytes::BytesMut;
    use socks5_proto::Address;
    use tokio::{
        io::{AsyncReadExt, AsyncWriteExt},
        net::TcpListener,
    };

    use super::TrojanConnection;
    use crate::proto::{
        padding::Padding,
        trojan::{Command, Request},
    };

    fn unique_hash() -> String {
        format!(
            "{:056x}",
            SystemTime::now()
                .duration_since(UNIX_EPOCH)
                .unwrap()
                .as_nanos()
        )
    }

    async fn local_listener() -> (TcpListener, SocketAddr) {
        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
        let addr = listener.local_addr().unwrap();
        (listener, addr)
    }

    #[tokio::test]
    async fn handle_connect_proxies_payload_in_both_directions() {
        let (listener, addr) = local_listener().await;
        let (server_side, mut client_side) = tokio::io::duplex(1024);
        let peer_addr = SocketAddr::from((Ipv4Addr::LOCALHOST, 40000));
        let hash = unique_hash();

        let handle_task = tokio::spawn(async move {
            TrojanConnection::new(server_side, "alice".into(), peer_addr)
                .handle()
                .await
        });
        let target_task = tokio::spawn(async move {
            let (mut target, _) = listener.accept().await.unwrap();
            let mut inbound = [0u8; 4];
            target.read_exact(&mut inbound).await.unwrap();
            target.write_all(b"pong").await.unwrap();
            target.shutdown().await.unwrap();
            inbound
        });

        let request = Request::new(hash, Command::Connect, Address::SocketAddress(addr));
        let mut raw = BytesMut::new();
        request.write_to_buf(&mut raw);
        raw.extend_from_slice(b"ping");

        client_side.write_all(&raw).await.unwrap();
        let mut response = [0u8; 4];
        client_side.read_exact(&mut response).await.unwrap();
        client_side.shutdown().await.unwrap();

        assert_eq!(&response, b"pong");
        assert_eq!(target_task.await.unwrap(), *b"ping");
        handle_task.await.unwrap().unwrap();
    }

    #[tokio::test]
    async fn handle_padding_request_replies_with_padding_then_proxies_payload() {
        let (listener, addr) = local_listener().await;
        let (server_side, mut client_side) = tokio::io::duplex(4096);
        let peer_addr = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 40001);
        let hash = unique_hash();

        let handle_task = tokio::spawn(async move {
            TrojanConnection::new(server_side, "alice".into(), peer_addr)
                .handle()
                .await
        });
        let target_task = tokio::spawn(async move {
            let (mut target, _) = listener.accept().await.unwrap();
            let mut inbound = [0u8; 5];
            target.read_exact(&mut inbound).await.unwrap();
            target.write_all(b"reply").await.unwrap();
            target.shutdown().await.unwrap();
            inbound
        });

        let request = Request::new(hash, Command::Padding, Address::SocketAddress(addr));
        let mut raw = BytesMut::new();
        request.write_to_buf(&mut raw);
        raw.extend_from_slice(b"hello");

        client_side.write_all(&raw).await.unwrap();
        let padding = Padding::read_from(&mut client_side).await.unwrap();
        assert!(padding.serialized_len() >= 258);

        let mut response = [0u8; 5];
        client_side.read_exact(&mut response).await.unwrap();
        client_side.shutdown().await.unwrap();

        assert_eq!(&response, b"reply");
        assert_eq!(target_task.await.unwrap(), *b"hello");
        handle_task.await.unwrap().unwrap();
    }
}