use std::net::SocketAddr;
use bytes::BytesMut;
use crate::domain::model::message_validation::validate_message_length;
#[derive(Clone)]
#[cfg(any(feature = "server", feature = "client"))]
pub struct InputBufVO {
data: BytesMut,
index: usize,
input_addr: Option<SocketAddr>,
}
#[cfg(any(feature = "server", feature = "client"))]
pub trait InputBufVOTrait {
fn get_constructor_id(&mut self) -> Option<u8>;
fn get_method_id(&mut self) -> Option<u16>;
fn next_u64(&mut self) -> Option<u64>;
fn next_u8(&mut self) -> Option<u8>;
fn next_str_with_len(&mut self, len: u64) -> Option<String>;
fn get_all_bytes(&self) -> BytesMut;
fn get_remaining_data_len(&self) -> usize;
}
impl InputBufVO {
pub(crate) fn new(buf: BytesMut, input_addr: SocketAddr) -> Self {
Self {
data: buf,
index: 3,
input_addr: Some(input_addr),
}
}
pub(crate) fn new_none() -> Self {
Self {
data: BytesMut::new(),
index: 3,
input_addr: None,
}
}
pub(crate) fn new_without_socket_addr(buf: BytesMut) -> Self {
Self {
data: buf,
index: 3,
input_addr: None,
}
}
pub fn get_input_addr(&self) -> Option<SocketAddr> {
self.input_addr
}
}
impl InputBufVOTrait for InputBufVO {
fn get_constructor_id(&mut self) -> Option<u8> {
let length = self.data.len();
if length < 1 {
None
} else {
let bytes = &self.data[0..1];
match bytes.try_into() {
Ok(value) => Some(u8::from_le_bytes(value)),
Err(_) => None,
}
}
}
fn get_method_id(&mut self) -> Option<u16> {
let length = self.data.len();
if length < 3 {
None
} else {
let bytes = &self.data[1..3];
match bytes.try_into() {
Ok(value) => Some(u16::from_le_bytes(value)),
Err(_) => None,
}
}
}
fn next_u64(&mut self) -> Option<u64> {
let length = self.data.len();
if length < self.index + 8 {
None
} else {
let bytes = &self.data[self.index..self.index + 8];
match bytes.try_into() {
Ok(value) => {
self.index += 8;
Some(u64::from_le_bytes(value))
},
Err(_) => None,
}
}
}
fn next_u8(&mut self) -> Option<u8> {
let length = self.data.len();
if length < self.index + 1 {
None
} else {
let bytes = &self.data[self.index..self.index + 1];
match bytes.try_into() {
Ok(value) => {
self.index += 1;
Some(u8::from_le_bytes(value))
},
Err(_) => None,
}
}
}
fn next_str_with_len(&mut self, len: u64) -> Option<String> {
let validated_len = match validate_message_length(len) {
Ok(l) => l,
Err(e) => {
tracing::warn!("Invalid string length in next_str_with_len: {}", e);
return None;
},
};
let length = self.data.len();
if length < self.index + validated_len {
None
} else {
let bytes = &mut self.data[self.index..self.index + validated_len].to_vec();
self.index += validated_len;
Some(String::from_utf8_lossy(bytes).to_string())
}
}
fn get_all_bytes(&self) -> BytesMut {
let mut vec = self.data.clone();
if vec.len() > 3 {
return vec.split_off(3);
} else {
vec = BytesMut::new();
}
vec
}
fn get_remaining_data_len(&self) -> usize {
let data_len = self.data.len();
let index = self.index;
if data_len == 0 {
0
} else {
data_len.saturating_sub(index)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::model::input_buf_vo::InputBufVOTrait;
fn sample_body() -> BytesMut {
let mut b = BytesMut::new();
b.extend_from_slice(&[1u8]);
b.extend_from_slice(&42u16.to_le_bytes());
b.extend_from_slice(b"hey");
b
}
#[test]
fn parses_fields_sequentially() {
let addr: SocketAddr = "127.0.0.1:1234".parse().unwrap();
let mut vo = InputBufVO::new(sample_body(), addr);
assert_eq!(vo.get_input_addr(), Some(addr));
assert_eq!(vo.get_constructor_id(), Some(1));
assert_eq!(vo.get_method_id(), Some(42));
assert_eq!(vo.next_str_with_len(3).as_deref(), Some("hey"));
assert_eq!(vo.get_remaining_data_len(), 0);
}
#[test]
fn reads_numbers_from_payload() {
let mut body = BytesMut::new();
body.extend_from_slice(&[1u8]);
body.extend_from_slice(&1u16.to_le_bytes());
body.extend_from_slice(&7u64.to_le_bytes());
body.extend_from_slice(&[9u8]);
let mut vo = InputBufVO::new(body, "127.0.0.1:1".parse().unwrap());
assert_eq!(vo.next_u64(), Some(7));
assert_eq!(vo.next_u8(), Some(9));
assert_eq!(vo.next_u64(), None);
assert_eq!(vo.next_u8(), None);
}
#[test]
fn short_buffers_return_none() {
let mut vo = InputBufVO::new_without_socket_addr(BytesMut::from(&[1u8, 2][..]));
assert_eq!(vo.get_constructor_id(), Some(1));
assert_eq!(vo.get_method_id(), None, "method id needs 3 bytes");
assert_eq!(vo.next_u64(), None);
assert_eq!(vo.next_str_with_len(1), None);
}
#[test]
fn oversized_string_len_is_rejected() {
let mut vo = InputBufVO::new(sample_body(), "127.0.0.1:1".parse().unwrap());
assert_eq!(vo.next_str_with_len(u64::MAX), None);
}
#[test]
fn get_all_bytes_strips_the_frame_prefix() {
let vo = InputBufVO::new(sample_body(), "127.0.0.1:1".parse().unwrap());
assert_eq!(&vo.get_all_bytes()[..], b"hey");
assert_eq!(&vo.get_all_bytes()[..], b"hey");
}
#[test]
fn empty_vo_behaves_gracefully() {
let mut vo = InputBufVO::new_none();
assert_eq!(vo.get_input_addr(), None);
assert_eq!(vo.get_constructor_id(), None);
assert_eq!(vo.get_method_id(), None);
assert!(vo.get_all_bytes().is_empty());
assert_eq!(vo.get_remaining_data_len(), 0);
}
#[test]
fn without_socket_addr_has_no_addr() {
let vo = InputBufVO::new_without_socket_addr(sample_body());
assert_eq!(vo.get_input_addr(), None);
assert_eq!(vo.get_remaining_data_len(), 3);
}
}