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>,
}
const UNSPECIFIED_IP: IpAddr = IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0));
pub struct TcpServer {
socket: net::TcpListener,
}
impl TcpServer {
pub async fn listen(params: &TcpConnectParams) -> Result<TcpServer, String> {
let (_, callback_port) = util::parse_address(¶ms.callback_addr)
.map_err(|err| format!("Invalid callback address: {err}"))?;
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 })
}
pub fn from_socket(socket: net::TcpListener) -> TcpServer {
TcpServer { socket }
}
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))
}
}
pub struct TcpClient {
controller_addr: String,
}
impl TcpClient {
pub fn new(controller_addr: &str) -> TcpClient {
TcpClient {
controller_addr: controller_addr.to_owned(),
}
}
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 {
src.reserve(total_needed - available);
Ok(None)
} else {
let data = src[HEADER_BYTES..total_needed].to_vec();
src.advance(total_needed);
Ok(Some(data))
}
}
}
pub struct TcpTransport {
stream: Framed<TcpStream, InsecureFrame>,
}
impl TcpTransport {
pub fn from(stream: TcpStream) -> TcpTransport {
TcpTransport {
stream: Framed::new(stream, InsecureFrame {}),
}
}
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::*;
#[tokio::test]
async fn network_frame() {
let mut writer = FramedWrite::new(Vec::<u8>::new(), InsecureFrame {});
writer.send("testing".as_bytes()).await.unwrap();
let expected = [
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());
}
}