use crate::{datapack::DataPack, error::ZerustError, request::Request, response::Response};
use std::net::SocketAddr;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpStream,
};
pub struct Connection {
stream: TcpStream,
pending_data: Vec<u8>,
}
impl Connection {
const HEADER_SIZE: usize = 8;
pub fn new(stream: TcpStream) -> Self {
Self {
stream,
pending_data: Vec::new(),
}
}
pub fn remote_addr(&self) -> Result<SocketAddr, ZerustError> {
self.stream.peer_addr().map_err(ZerustError::IoError)
}
pub async fn read_request(&mut self) -> Result<Request, ZerustError> {
let header_bytes = self.read_exact(Self::HEADER_SIZE).await?;
let (msg_id, data_len) = DataPack::unpack_header(&header_bytes)?;
let data = if data_len > 0 {
self.read_exact(data_len as usize).await?
} else {
Vec::new()
};
Ok(Request::new(msg_id, data))
}
async fn read_exact(&mut self, size: usize) -> Result<Vec<u8>, ZerustError> {
while self.pending_data.len() < size {
let mut buffer = [0u8; 1024]; let n = self.stream.read(&mut buffer).await?;
if n == 0 {
return Err(ZerustError::ConnectionClosed);
}
self.pending_data.extend_from_slice(&buffer[..n]);
}
let result = self.pending_data.drain(..size).collect(); Ok(result)
}
pub async fn send_response(&mut self, resp: &Response) -> Result<(), ZerustError> {
let bytes = DataPack::pack(resp.msg_id(), resp.data());
self.stream.write_all(&bytes).await?;
Ok(())
}
}