From 6a62c33f4283de714894b5b8a582d9baf0a001e6 Mon Sep 17 00:00:00 2001 From: cool-mist Date: Thu, 3 Sep 2026 22:47:50 +0530 Subject: [PATCH] WIP --- crates/hd-lib/src/bus/mod.rs | 2 +- crates/hd-lib/src/bus/packet.rs | 268 +++++++++++++++++--------------- crates/hd-lib/src/error/mod.rs | 36 ++++- crates/hd-server/src/lib.rs | 147 +++++++++++------- 4 files changed, 267 insertions(+), 186 deletions(-) diff --git a/crates/hd-lib/src/bus/mod.rs b/crates/hd-lib/src/bus/mod.rs index c1b1301..4163293 100644 --- a/crates/hd-lib/src/bus/mod.rs +++ b/crates/hd-lib/src/bus/mod.rs @@ -1,4 +1,4 @@ -mod packet; +pub mod packet; use tracing::{info, info_span}; diff --git a/crates/hd-lib/src/bus/packet.rs b/crates/hd-lib/src/bus/packet.rs index ce6d249..15bfc7a 100644 --- a/crates/hd-lib/src/bus/packet.rs +++ b/crates/hd-lib/src/bus/packet.rs @@ -1,12 +1,15 @@ -use std::{ - io::{Read, Write}, - net::TcpStream, -}; +use std::{fmt::Display, sync::Arc}; -use crate::bus::EventData; +use crate::bus::{Event, EventData}; -enum PacketType { +mod packet_constants { + pub(crate) const HEADER_SIZE: usize = 3; +} + +#[derive(PartialEq)] +pub enum PacketType { Connect, + Subscribe, Unsubscribe, Publish, SendEvent, @@ -15,8 +18,9 @@ enum PacketType { Disconnect, } -enum Packet { - Connect(ConnectPacket), +pub enum Packet { + Connect, + Subscribe(SubscribePacket), Unsubscribe, Publish(PublishPacket), SendEvent(SendEventPacket), @@ -25,186 +29,181 @@ enum Packet { Disconnect, } -struct ConnectPacket { +pub struct SubscribePacket { cursor: u64, } -struct PublishPacket { - event_type: u8, - data_utf8: Vec, +pub struct PublishPacket { + pub event_type: u8, + pub data_utf8: Vec, } -struct SendEventPacket { +pub struct SendEventPacket { event_id: u64, event_type: u8, data_utf8: Vec, } -struct SettlePacket { +pub struct SettlePacket { cursor: u64, } -struct AckPacket { +pub struct AckPacket { packet_type: PacketType, } impl Packet { - fn get_packet_type(&self) -> PacketType { - match self { - Packet::Connect(_) => PacketType::Connect, - Packet::Unsubscribe => PacketType::Unsubscribe, - Packet::Publish(_) => PacketType::Publish, - Packet::SendEvent(_) => PacketType::SendEvent, - Packet::Settle(_) => PacketType::Settle, - Packet::Ack(_) => PacketType::Ack, - Packet::Disconnect => PacketType::Disconnect, - } + pub fn create_connect_packet(cursor: u64) -> Packet { + Packet::Subscribe(SubscribePacket { cursor }) } - fn create_connect_packet(cursor: u64) -> Packet { - Packet::Connect(ConnectPacket { cursor }) - } - - fn create_unsubscribe_packet() -> Packet { + pub fn create_unsubscribe_packet() -> Packet { Packet::Unsubscribe } - fn create_publish_packet(evt: EventData) -> Packet { + // TODO: Avoid copying data_utf8 + pub fn create_publish_packet(evt: EventData) -> Packet { Packet::Publish(PublishPacket { event_type: evt.event_type, data_utf8: evt.data_utf8, }) } - fn create_settle_packet(cursor: u64) -> Packet { + // TODO: Avoid copying data_utf8 + pub fn create_send_event_packet(evt: Arc) -> Packet { + Packet::SendEvent(SendEventPacket { + event_id: evt.id, + event_type: evt.data.event_type, + data_utf8: evt.data.data_utf8.clone(), + }) + } + + pub fn create_settle_packet(cursor: u64) -> Packet { Packet::Settle(SettlePacket { cursor }) } - fn create_ack_packet(packet_type: PacketType) -> Packet { + pub fn create_ack_packet(packet_type: PacketType) -> Packet { Packet::Ack(AckPacket { packet_type }) } - fn create_disconnect_packet() -> Packet { + pub fn create_disconnect_packet() -> Packet { Packet::Disconnect } } -struct PacketFrame { - header: HeaderFrame, - data: DataFrame, -} - -/// 3 Bytes -struct HeaderFrame { - /// 0 Connect - /// 1 Subscribe - /// 2 Unsubscribe - /// 3 Publish - /// 4 Settle - /// 5 Ack - /// 6 Disconnect - packet_type: u8, - /// Length of data frame, Little Endian, so 5 = [5 0], 256 = [255 1] - data_frame_length: [u8; 2], -} - -struct DataFrame { - data: Vec, -} - -enum PacketError { +#[derive(Debug)] +pub enum PacketError { WrongPacketType(u8), PacketProtocolError, - TcpStreamError(std::io::Error), } -fn read_packet(stream: &mut TcpStream) -> Result { - let mut header_buf = vec![0; 3]; - stream - .read_exact(&mut header_buf) - .map_err(PacketError::TcpStreamError)?; +impl Display for PacketError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let to_write = match self { + PacketError::WrongPacketType(p) => &format!("WRONG_PACKET_TYPE({})", p), + PacketError::PacketProtocolError => "PACKET_PROTOCOL_ERROR", + }; - let packet_type = get_packet_type_from_byte(header_buf[0])?; + write!(f, "{}", to_write) + } +} - match packet_type { - PacketType::Connect => { - let cursor = read_u64(stream)?; - Ok(Packet::Connect(ConnectPacket { cursor })) +pub struct ReadPacketState { + pub required_buffer_size: usize, + packet_type: PacketType, +} + +pub fn start_read_packet(buf: &mut [u8; packet_constants::HEADER_SIZE]) -> Result { + let packet_type = get_packet_type_from_byte(buf[0])?; + let required_buffer_size = read_u16(&buf[1..3]); + + Ok(ReadPacketState { + required_buffer_size, + packet_type, + }) +} + +pub fn read_packet(buf: &[u8], read_state: ReadPacketState) -> Result { + match read_state.packet_type { + PacketType::Connect => Ok(Packet::Connect), + PacketType::Subscribe => { + let cursor = read_u64(&buf[0..])?; + Ok(Packet::Subscribe(SubscribePacket { cursor })) } PacketType::Unsubscribe => Ok(Packet::Unsubscribe), PacketType::Publish => { - let mut header = vec![0u8; 3]; - stream.read_exact(&mut header).map_err(PacketError::TcpStreamError)?; + let event_type = buf[0]; + let len = read_u16(&buf[1..]); - let event_type = header[0]; - - let len = usize::from_le_bytes([header[1], header[2], 0, 0, 0, 0, 0, 0]); - let mut data_utf8 = vec![0u8; len]; - stream.read_exact(&mut data_utf8).map_err(PacketError::TcpStreamError)?; - - let publish_packet = PublishPacket { event_type, data_utf8 }; + // TODO: Share underlying bytes without copying + let publish_packet = PublishPacket { + event_type, + data_utf8: buf[3..len].to_vec(), + }; Ok(Packet::Publish(publish_packet)) } PacketType::SendEvent => { - let mut header = vec![0u8; 11]; - stream.read_exact(&mut header).map_err(PacketError::TcpStreamError)?; - - let mut le_bytes = [0u8; 8]; - le_bytes.copy_from_slice(&header[0..8]); - let event_id = u64::from_le_bytes(le_bytes); - - let event_type = header[8]; - - let len = usize::from_le_bytes([header[9], header[10], 0, 0, 0, 0, 0, 0]); - let mut data_utf8 = vec![0u8; len]; - stream.read_exact(&mut data_utf8).map_err(PacketError::TcpStreamError)?; + let event_id = read_u64(&buf)?; + let event_type = buf[8]; + let len = usize::from_le_bytes([buf[9], buf[10], 0, 0, 0, 0, 0, 0]); let send_event_packet = SendEventPacket { event_id, event_type, - data_utf8, + data_utf8: buf[11..len].to_vec(), }; Ok(Packet::SendEvent(send_event_packet)) } PacketType::Settle => { - let cursor = read_u64(stream)?; + let mut le_bytes = [0u8; 8]; + le_bytes.copy_from_slice(&buf[0..8]); + let cursor = u64::from_le_bytes(le_bytes); Ok(Packet::Settle(SettlePacket { cursor })) } PacketType::Ack => { - let packet_type = get_packet_type_from_byte(read_byte(stream)?)?; + let packet_type = get_packet_type_from_byte(buf[0])?; Ok(Packet::Ack(AckPacket { packet_type })) } PacketType::Disconnect => Ok(Packet::Disconnect), } } -fn write_packet(stream: &mut TcpStream, packet: &Packet) -> Result<(), PacketError> { - let packet_type = packet.get_packet_type(); +pub struct WritePacketState { + pub required_buffer_size: usize, + data_frame_length: usize, + packet_type: PacketType, +} + +pub fn start_write_packet(packet: &Packet) -> WritePacketState { let data_frame_length = calculate_data_frame_length(&packet); + let required_buffer_size = packet_constants::HEADER_SIZE + data_frame_length; + let packet_type = get_packet_type(packet); + WritePacketState { + required_buffer_size, + data_frame_length, + packet_type, + } +} - let mut buf = vec![0u8; 3 + data_frame_length]; - write_header_frame(&mut buf, packet_type, data_frame_length); - - if data_frame_length > 0 { +pub fn write_packet(buf: &mut [u8], packet: &Packet, write_state: WritePacketState) -> Result<(), PacketError> { + write_header_frame(buf, write_state.packet_type, write_state.data_frame_length); + if write_state.data_frame_length > 0 { write_data_frame(&mut buf[3..], packet); } - stream.write_all(&buf).map_err(PacketError::TcpStreamError)?; - stream.flush().map_err(PacketError::TcpStreamError)?; - Ok(()) } fn write_header_frame(header_frame: &mut [u8], packet_type: PacketType, data_frame_length: usize) { header_frame[0] = get_byte_from_packet_type(&packet_type); if data_frame_length > 0 { - write_length(data_frame_length, &mut header_frame[1..3]); + write_u16(data_frame_length, &mut header_frame[1..3]); } } fn write_data_frame(data_frame: &mut [u8], packet: &Packet) { match packet { - Packet::Connect(connect_packet) => { + Packet::Subscribe(connect_packet) => { write_u64(connect_packet.cursor, data_frame); } Packet::Publish(publish_packet) => { @@ -230,7 +229,7 @@ fn write_data_frame(data_frame: &mut [u8], packet: &Packet) { fn write_event_data(event_type: u8, data_utf8: &[u8], buf: &mut [u8]) { buf[0] = event_type; - write_length(data_utf8.len(), &mut buf[1..3]); + write_u16(data_utf8.len(), &mut buf[1..3]); let available_length = buf.len() - 3; buf[3..].copy_from_slice(&data_utf8[..available_length]); } @@ -240,33 +239,54 @@ fn write_u64(v64: u64, data_frame: &mut [u8]) { data_frame[0..8].copy_from_slice(&bytes); } -fn read_u64(stream: &mut TcpStream) -> Result { +fn write_u16(mut length: usize, buf: &mut [u8]) { + if length > 0xffff { + length = 0xffff; + } + + let bytes = length.to_le_bytes(); + buf[0..2].copy_from_slice(&bytes); +} + +fn read_u64(buf: &[u8]) -> Result { let mut le_bytes = [0u8; 8]; - stream.read_exact(&mut le_bytes).map_err(PacketError::TcpStreamError)?; + le_bytes.copy_from_slice(buf); Ok(u64::from_le_bytes(le_bytes)) } -fn read_byte(stream: &mut TcpStream) -> Result { - let mut byte = [0u8; 1]; - stream.read_exact(&mut byte).map_err(PacketError::TcpStreamError)?; - Ok(byte[0]) +fn read_u16(buf: &[u8]) -> usize { + usize::from_le_bytes([buf[0], buf[1], 0, 0, 0, 0, 0, 0]) +} + +pub fn get_packet_type(packet: &Packet) -> PacketType { + match packet { + Packet::Connect => PacketType::Connect, + Packet::Subscribe(_) => PacketType::Subscribe, + Packet::Unsubscribe => PacketType::Unsubscribe, + Packet::Publish(_) => PacketType::Publish, + Packet::SendEvent(_) => PacketType::SendEvent, + Packet::Settle(_) => PacketType::Settle, + Packet::Ack(_) => PacketType::Ack, + Packet::Disconnect => PacketType::Disconnect, + } } fn get_byte_from_packet_type(packet_type: &PacketType) -> u8 { match packet_type { PacketType::Connect => 0, - PacketType::Unsubscribe => 1, - PacketType::Publish => 2, - PacketType::SendEvent => 3, - PacketType::Settle => 4, - PacketType::Ack => 5, - PacketType::Disconnect => 6, + PacketType::Subscribe => 1, + PacketType::Unsubscribe => 2, + PacketType::Publish => 3, + PacketType::SendEvent => 4, + PacketType::Settle => 5, + PacketType::Ack => 6, + PacketType::Disconnect => 7, } } fn get_packet_type_from_byte(byte: u8) -> Result { let packet_type = match byte { - 0 => Some(PacketType::Connect), + 0 => Some(PacketType::Subscribe), 1 => Some(PacketType::Unsubscribe), 2 => Some(PacketType::Publish), 3 => Some(PacketType::SendEvent), @@ -283,18 +303,10 @@ fn get_packet_type_from_byte(byte: u8) -> Result { Err(PacketError::WrongPacketType(byte)) } -fn write_length(mut length: usize, buf: &mut [u8]) { - if length > 0xffff { - length = 0xffff; - } - - let bytes = length.to_le_bytes(); - buf[0..2].copy_from_slice(&bytes); -} - fn calculate_data_frame_length(packet: &Packet) -> usize { let mut ret = match packet { - Packet::Connect(_) => 8, + Packet::Connect => 0, + Packet::Subscribe(_) => 8, Packet::Unsubscribe => 0, Packet::Publish(publish_packet) => 1 + 2 + publish_packet.data_utf8.len(), Packet::SendEvent(send_event_packet) => 8 + 1 + 2 + send_event_packet.data_utf8.len(), diff --git a/crates/hd-lib/src/error/mod.rs b/crates/hd-lib/src/error/mod.rs index fa98d0e..5d42cf1 100644 --- a/crates/hd-lib/src/error/mod.rs +++ b/crates/hd-lib/src/error/mod.rs @@ -1,12 +1,17 @@ -use std::fmt::Display; +use std::{fmt::Display, string::FromUtf8Error}; + +use crate::bus::packet::PacketError; #[derive(Debug)] pub enum HError { BusLockPoisoned(String), SubscriptionNotFound(String), - TcpSocketBindError(std::io::Error), - TcpPeerError(std::io::Error), + IoError(std::io::Error), + TcpPacketError(PacketError), + Utf8FormatError(FromUtf8Error), NoMoreEvents, + ProtocolError, + PeerDisconnect, } impl Display for HError { @@ -14,9 +19,30 @@ impl Display for HError { match self { HError::BusLockPoisoned(s) => write!(f, "{}", s), HError::SubscriptionNotFound(s) => write!(f, "{}", s), - HError::TcpSocketBindError(error) => write!(f, "{}", error), - HError::TcpPeerError(error) => write!(f, "{}", error), + HError::IoError(error) => write!(f, "{}", error), HError::NoMoreEvents => write!(f, "{}", "No more events in the bus for now"), + HError::TcpPacketError(packet_error) => write!(f, "Failed to read/write packets due to {}", packet_error), + HError::PeerDisconnect => write!(f, "Peer disconnect"), + HError::Utf8FormatError(e) => write!(f, "Utf8 formatting error {}", e), + HError::ProtocolError => write!(f, "Protocol error"), } } } + +impl From for HError { + fn from(value: std::io::Error) -> Self { + HError::IoError(value) + } +} + +impl From for HError { + fn from(value: PacketError) -> Self { + HError::TcpPacketError(value) + } +} + +impl From for HError { + fn from(value: FromUtf8Error) -> Self { + HError::Utf8FormatError(value) + } +} diff --git a/crates/hd-server/src/lib.rs b/crates/hd-server/src/lib.rs index 608c321..d6a69e6 100644 --- a/crates/hd-server/src/lib.rs +++ b/crates/hd-server/src/lib.rs @@ -6,7 +6,12 @@ use std::{ }; use hd_lib::{ - bus::{Event, EventBus, EventData}, error::HError, thread_pool::ThreadPool, + bus::{ + Event, EventBus, EventData, + packet::{self, Packet, PacketType}, + }, + error::HError, + thread_pool::ThreadPool, }; use tracing::{Span, info, info_span}; @@ -34,7 +39,7 @@ impl TcpEventBus { impl TcpEventBus { pub fn start(&self, addr: &'static str) -> Result<(), HError> { - let listener = TcpListener::bind(addr).map_err(HError::TcpSocketBindError)?; + let listener = TcpListener::bind(addr)?; info!("Heimdall listening at {}", addr); for stream in listener.incoming() { match stream { @@ -50,16 +55,15 @@ impl TcpEventBus { #[derive(Copy, Clone, PartialEq)] enum EventBusClientConnectionState { Connect, - ReadPrelude, - ReadCursor, + Command, SendEvent, - ReceiveAck, + WaitEvent, Disconnect, } struct EventBusClientConnection { - stream: TcpStream, pub addr: String, + stream: TcpStream, state: EventBusClientConnectionState, subscription: Option>>, bus: EventBus, @@ -72,9 +76,7 @@ impl EventBusClientConnection { .map(|a| a.to_string()) .unwrap_or("unknown".to_string()); - stream - .set_read_timeout(Some(Duration::from_secs(5))) - .map_err(HError::TcpPeerError)?; + stream.set_read_timeout(Some(Duration::from_secs(5)))?; let state = EventBusClientConnectionState::Connect; @@ -88,6 +90,47 @@ impl EventBusClientConnection { } } +fn read_ack(stream: &mut TcpStream, packet_type: PacketType) -> Result<(), HError> { + let packet = read_packet(stream)?; + if packet_type == packet::get_packet_type(&packet) { + return Ok(()) + } + + Err(HError::ProtocolError) +} + +fn read_packet(stream: &mut TcpStream) -> Result { + let mut header_buf = [0u8; 3]; + stream.read_exact(&mut header_buf)?; + let read_state = packet::start_read_packet(&mut header_buf)?; + + let mut buf = vec![0u8; read_state.required_buffer_size]; + stream.read_exact(&mut buf)?; + let packet = packet::read_packet(&buf[0..], read_state)?; + + if let Packet::Disconnect = packet { + return Err(HError::PeerDisconnect); + }; + + Ok(packet) +} + +fn write_ack(stream: &mut TcpStream, packet_type: PacketType) -> Result<(), HError> { + let packet = Packet::create_ack_packet(packet_type); + write_packet(stream, &packet) +} + +fn write_packet(stream: &mut TcpStream, packet: &Packet) -> Result<(), HError> { + let write_state = packet::start_write_packet(packet); + let mut buf = vec![0u8; write_state.required_buffer_size]; + packet::write_packet(&mut buf, packet, write_state)?; + + stream.write_all(&mut buf)?; + stream.flush()?; + + Ok(()) +} + pub fn handle_event_bus_client(thread_pool: &ThreadPool, bus: EventBus, stream: TcpStream) -> Result<(), HError> { thread_pool.execute(move || { let mut client = EventBusClientConnection::new(bus, stream)?; @@ -120,11 +163,10 @@ pub fn handle_event_bus_client(thread_pool: &ThreadPool, bus: EventBus, stream: fn handle_event_bus_client_state(client: &mut EventBusClientConnection) -> Result<(), HError> { let next_state = match client.state { EventBusClientConnectionState::Connect => handle_state_connect(client), - EventBusClientConnectionState::ReadPrelude => todo!(), - EventBusClientConnectionState::ReadCursor => todo!(), + EventBusClientConnectionState::Command => handle_state_connect(client), EventBusClientConnectionState::SendEvent => handle_state_send_event(client), - EventBusClientConnectionState::ReceiveAck => todo!(), - _ => Ok(EventBusClientConnectionState::Disconnect), + EventBusClientConnectionState::WaitEvent => handle_state_wait_event(client), + EventBusClientConnectionState::Disconnect => Ok(EventBusClientConnectionState::Disconnect), }?; client.state = next_state; @@ -132,15 +174,34 @@ fn handle_event_bus_client_state(client: &mut EventBusClientConnection) -> Resul } fn handle_state_connect(client: &mut EventBusClientConnection) -> Result { - let mut buf = [0u8; 1]; - client.stream.read_exact(&mut buf).map_err(HError::TcpPeerError)?; - if buf[0] == 1 { - let subscription = client.bus.subscribe()?; - client.subscription = Some(subscription); - return Ok(EventBusClientConnectionState::SendEvent); - } + let Packet::Connect = read_packet(&mut client.stream)? else { + return Ok(EventBusClientConnectionState::Disconnect); + }; - Ok(EventBusClientConnectionState::Disconnect) + write_ack(&mut client.stream, PacketType::Connect)?; + Ok(EventBusClientConnectionState::Command) +} + +fn handle_state_command(client: &mut EventBusClientConnection) -> Result { + match read_packet(&mut client.stream)? { + // TODO: Read cursor + Packet::Subscribe(_) => { + write_ack(&mut client.stream, PacketType::Subscribe)?; + + let subscription = client.bus.subscribe()?; + client.subscription = Some(subscription); + Ok(EventBusClientConnectionState::SendEvent) + } + Packet::Publish(publish_packet) => { + write_ack(&mut client.stream, PacketType::Publish)?; + + // TODO: Possible without copying data_utf8? + let data = String::from_utf8(publish_packet.data_utf8)?; + client.bus.publish(EventData::new(publish_packet.event_type, &data))?; + Ok(EventBusClientConnectionState::WaitEvent) + } + _ => Ok(EventBusClientConnectionState::Disconnect), + } } fn handle_state_send_event(client: &mut EventBusClientConnection) -> Result { @@ -150,56 +211,38 @@ fn handle_state_send_event(client: &mut EventBusClientConnection) -> Result ID: {} TYPE: {} LEN: {} <=", - evt.id, evt.data.event_type, data_len + evt.id, + evt.data.event_type, + evt.data.data_utf8.len() ); - let stream = &mut client.stream; - stream.write(&mut buf).map_err(HError::TcpPeerError)?; - stream.flush().map_err(HError::TcpPeerError)?; - - let mut ack = vec![0]; - stream.read(&mut ack).map_err(HError::TcpPeerError)?; + let packet = Packet::create_send_event_packet(evt); + write_packet(&mut client.stream, &packet)?; + read_ack(&mut client.stream, PacketType::SendEvent)?; } } Ok(EventBusClientConnectionState::Disconnect) } +fn handle_state_wait_event(client: &mut EventBusClientConnection) -> Result { + todo!() +} + fn write_le_bytes(v64: u64, buf: &mut [u8]) { let le_bytes = v64.to_le_bytes(); - buf[0] = le_bytes[0]; - buf[1] = le_bytes[1]; - buf[2] = le_bytes[2]; - buf[3] = le_bytes[3]; - buf[4] = le_bytes[4]; - buf[5] = le_bytes[5]; - buf[6] = le_bytes[6]; - buf[7] = le_bytes[7]; + buf[0..8].copy_from_slice(&le_bytes); } impl std::fmt::Display for EventBusClientConnectionState { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { let to_write = match self { EventBusClientConnectionState::Connect => "CONNECT", - EventBusClientConnectionState::ReadPrelude => "READ_PRELUDE", - EventBusClientConnectionState::ReadCursor => "READ_CURSOR", + EventBusClientConnectionState::Command => "COMMAND", EventBusClientConnectionState::SendEvent => "SEND_EVENT", - EventBusClientConnectionState::ReceiveAck => "RECEIVE_ACK", + EventBusClientConnectionState::WaitEvent => "WAIT_EVENT", EventBusClientConnectionState::Disconnect => "DISCONNECT", };