breakmancer 0.9.0

Drop a breakpoint into any shell.
Documentation
//! Implementation of TCP transport channel.
//!
//! This is to help avoid mingling transport-specific logic with other
//! protocol code.

use std::{
    io::{Error, ErrorKind},
    net::{IpAddr, Ipv6Addr, SocketAddr, ToSocketAddrs},
};

use futures_util::sink::SinkExt;
use tokio::net::{self, TcpStream};
use tokio_stream::StreamExt;
use tokio_util::{
    bytes::{Buf, BytesMut},
    codec::{Decoder, Encoder, Framed},
};

use crate::util;

#[derive(Clone, Debug)]
pub struct TcpConnectParams {
    pub callback_addr: String,
    pub local_port: Option<u16>,
}

// Should work for listening on both IPv4 and IPv6.
const UNSPECIFIED_IP: IpAddr = IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0));

/// TCP server, listening on a socket, but not yet connected.
pub struct TcpServer {
    socket: net::TcpListener,
}

impl TcpServer {
    /// Create a listener by binding to a port.
    pub async fn listen(params: &TcpConnectParams) -> Result<TcpServer, String> {
        // We can't really validate the callback address beyond checking
        // if it's well-formed; it's possible that the breakpoint needs to
        // be given a domain name that won't resolve from the controller's
        // vantage point.
        let (_, callback_port) = util::parse_address(&params.callback_addr)
            .map_err(|err| format!("Invalid callback address: {err}"))?;

        // We'll still use it if it doesn't resolve, but at least warn the user.
        match params.callback_addr.to_socket_addrs() {
            Ok(addrs) => {
                if addrs.len() == 0 {
                    eprintln!("[WARNING] Callback host did not resolve to any IP addresses.");
                }
            }
            Err(err) => {
                eprintln!("[WARNING] Could not resolve callback host to an IP address: {err}")
            }
        }
        let listen_port = params.local_port.unwrap_or(callback_port);

        let listen_addr = SocketAddr::new(UNSPECIFIED_IP, listen_port);
        let socket = net::TcpListener::bind(listen_addr)
            .await
            .map_err(|err| format!("Could not listen on {listen_addr}: {err}"))?;

        Ok(TcpServer { socket })
    }

    /// Utility for integration tests to use.
    ///
    /// Just directly wraps an existing server socket.
    pub fn from_socket(socket: net::TcpListener) -> TcpServer {
        TcpServer { socket }
    }

    // Transport interface

    pub fn describe_listening(&self) -> String {
        format!("listening on {}", self.socket.local_addr().unwrap())
    }

    pub async fn accept_connection(&mut self) -> Result<TcpTransport, String> {
        let (stream, client_addr) = self
            .socket
            .accept()
            .await
            .map_err(|err| format!("Couldn't get new connection: {err}"))?;
        println!("Connection from {client_addr:?}");

        Ok(TcpTransport::from(stream))
    }
}

/// A TCP client, able to make an outbound connection.
pub struct TcpClient {
    controller_addr: String,
}

impl TcpClient {
    /// Create from the address of the controller.
    pub fn new(controller_addr: &str) -> TcpClient {
        TcpClient {
            controller_addr: controller_addr.to_owned(),
        }
    }

    // Transport interface

    pub async fn new_connection(&self) -> Result<TcpTransport, String> {
        let stream = net::TcpStream::connect(&self.controller_addr)
            .await
            .map_err(|err| format!("Unable to connect to {}: {err}", &self.controller_addr))?;
        match stream.peer_addr() {
            Ok(connected) => eprintln!("Connecting to controller at {connected}"),
            Err(err) => eprintln!("Unable to get address of controller connection: {err}"),
        }

        Ok(TcpTransport::from(stream))
    }
}

const HEADER_BYTES: usize = size_of::<u32>();
const PAYLOAD_MAX: usize = 100 * 1024 * 1024;

struct InsecureFrame {}

impl Encoder<&[u8]> for InsecureFrame {
    type Error = std::io::Error;

    fn encode(&mut self, item: &[u8], dst: &mut BytesMut) -> Result<(), Self::Error> {
        let data_len = item.len();
        if data_len > PAYLOAD_MAX {
            return Err(Error::new(
                ErrorKind::InvalidData,
                format!("Cannot send message frame carrying more than {PAYLOAD_MAX} bytes"),
            ));
        }
        let len_bytes = u32::to_le_bytes(data_len as u32);

        dst.reserve(len_bytes.len() + data_len);
        dst.extend_from_slice(&len_bytes);
        dst.extend_from_slice(item);

        Ok(())
    }
}

impl Decoder for InsecureFrame {
    type Item = Vec<u8>;
    type Error = std::io::Error;

    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
        if src.len() < HEADER_BYTES {
            return Ok(None);
        }

        let mut len_header = [0u8; HEADER_BYTES];
        len_header.copy_from_slice(&src[..HEADER_BYTES]);
        let length = u32::from_le_bytes(len_header) as usize;

        let total_needed = HEADER_BYTES + length;
        let available = src.len();
        if total_needed > available {
            // Performance enhancement -- prep for full frame.
            src.reserve(total_needed - available);
            Ok(None)
        } else {
            let data = src[HEADER_BYTES..total_needed].to_vec();
            src.advance(total_needed);
            Ok(Some(data))
        }
    }
}

/// An open TCP connection.
///
/// TCP connection is closed when stream is dropped.
pub struct TcpTransport {
    stream: Framed<TcpStream, InsecureFrame>,
}

impl TcpTransport {
    pub fn from(stream: TcpStream) -> TcpTransport {
        TcpTransport {
            stream: Framed::new(stream, InsecureFrame {}),
        }
    }

    // Transport interface

    pub async fn send_raw(&mut self, raw: &[u8]) -> Result<(), String> {
        self.stream.send(raw).await.map_err(|err| format!("{err}"))
    }

    pub async fn receive_raw(&mut self) -> Result<Option<Vec<u8>>, String> {
        self.stream
            .next()
            .await
            .transpose()
            .map_err(|err| format!("{err}"))
    }
}

#[cfg(test)]
mod tests {
    use tokio_stream::StreamExt;
    use tokio_util::codec::{FramedRead, FramedWrite};

    use super::*;

    /// Pinning test for network frame format.
    #[tokio::test]
    async fn network_frame() {
        let mut writer = FramedWrite::new(Vec::<u8>::new(), InsecureFrame {});
        writer.send("testing".as_bytes()).await.unwrap();

        let expected = [
            // Length header, 7 byte body.
            b"\x07\x00\x00\x00" as &[u8],
            b"testing",
        ]
        .concat();
        assert_eq!(
            expected.escape_ascii().to_string(),
            writer.get_ref().escape_ascii().to_string()
        );

        let mut reader = FramedRead::new(writer.get_ref().as_slice(), InsecureFrame {});
        let got_bytes = reader.next().await.unwrap().unwrap();
        assert_eq!(got_bytes.as_slice(), b"testing");
    }

    #[tokio::test]
    async fn decode_with_partial_header() {
        let source = "\x07\x00".as_bytes();
        let mut reader = FramedRead::new(source, InsecureFrame {});
        assert!(format!("{:?}", reader.next().await).contains("bytes remaining on stream"),);
        assert_eq!(reader.read_buffer().to_vec(), b"\x07\x00".to_vec());
    }

    #[tokio::test]
    async fn decode_with_partial_body() {
        let source = "\x07\x00\x00\x00_".as_bytes();
        let mut reader = FramedRead::new(source, InsecureFrame {});
        assert!(format!("{:?}", reader.next().await).contains("bytes remaining on stream"),);
        assert_eq!(reader.read_buffer().to_vec(), b"\x07\x00\x00\x00_".to_vec());
    }

    #[tokio::test]
    async fn decode_with_surplus() {
        let source = "\x07\x00\x00\x00testingMORE".as_bytes();
        let mut reader = FramedRead::new(source, InsecureFrame {});
        assert_eq!(reader.next().await.unwrap().unwrap(), b"testing");
        assert_eq!(reader.read_buffer().to_vec(), b"MORE".to_vec());
    }
}