use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use uuid::Uuid;
use super::{Decode, DecodeError, Encode, EncodeError, take};
impl Encode for u8 {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
out.push(*self);
Ok(())
}
}
impl Decode for u8 {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
Ok(take(input, 1)?[0])
}
}
impl Encode for u32 {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
out.extend_from_slice(&self.to_le_bytes());
Ok(())
}
}
impl Decode for u32 {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
let bytes = take(input, 4)?;
Ok(u32::from_le_bytes(
bytes.try_into().expect("length checked"),
))
}
}
impl Encode for u64 {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
out.extend_from_slice(&self.to_le_bytes());
Ok(())
}
}
impl Decode for u64 {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
let bytes = take(input, 8)?;
Ok(u64::from_le_bytes(
bytes.try_into().expect("length checked"),
))
}
}
impl Encode for Uuid {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
out.extend_from_slice(self.as_bytes());
Ok(())
}
}
impl Decode for Uuid {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
let bytes = take(input, 16)?;
Ok(Uuid::from_bytes(bytes.try_into().expect("length checked")))
}
}
impl Encode for String {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
let length: u8 = self
.len()
.try_into()
.map_err(|_| EncodeError::StringTooLong { length: self.len() })?;
out.push(length);
out.extend_from_slice(self.as_bytes());
Ok(())
}
}
impl Decode for String {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
let length = u8::decode(input)? as usize;
let bytes = take(input, length)?;
Ok(String::from_utf8(bytes.to_vec())?)
}
}
impl Encode for IpAddr {
fn encode(&self, out: &mut Vec<u8>) -> Result<(), EncodeError> {
match self {
IpAddr::V4(address) => {
out.push(4);
out.extend_from_slice(&address.octets());
}
IpAddr::V6(address) => {
out.push(6);
out.extend_from_slice(&address.octets());
}
}
Ok(())
}
}
impl Decode for IpAddr {
fn decode(input: &mut &[u8]) -> Result<Self, DecodeError> {
let version = u8::decode(input)?;
match version {
4 => {
let octets: [u8; 4] = take(input, 4)?.try_into().expect("length checked");
Ok(IpAddr::V4(Ipv4Addr::from(octets)))
}
6 => {
let octets: [u8; 16] = take(input, 16)?.try_into().expect("length checked");
Ok(IpAddr::V6(Ipv6Addr::from(octets)))
}
other => Err(DecodeError::InvalidIpVersion(other)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::codec::{decode_payload, encode_payload};
#[test]
fn integers_use_little_endian() {
assert_eq!(
encode_payload(&0x1234_5678u32).unwrap(),
[0x78, 0x56, 0x34, 0x12]
);
assert_eq!(
encode_payload(&0x0102_0304_0506_0708u64).unwrap(),
[0x08, 0x07, 0x06, 0x05, 0x04, 0x03, 0x02, 0x01]
);
assert_eq!(
decode_payload::<u32>(&[0x78, 0x56, 0x34, 0x12]).unwrap(),
0x1234_5678
);
}
#[test]
fn string_is_u8_length_prefixed() {
assert_eq!(
encode_payload(&"foo".to_string()).unwrap(),
[3, b'f', b'o', b'o']
);
assert_eq!(
decode_payload::<String>(&[3, b'f', b'o', b'o']).unwrap(),
"foo"
);
assert_eq!(decode_payload::<String>(&[0]).unwrap(), "");
}
#[test]
fn overlong_string_fails_to_encode() {
let error = encode_payload(&"x".repeat(256)).unwrap_err();
assert!(matches!(error, EncodeError::StringTooLong { length: 256 }));
}
#[test]
fn invalid_utf8_fails_to_decode() {
let error = decode_payload::<String>(&[2, 0xff, 0xff]).unwrap_err();
assert!(matches!(error, DecodeError::InvalidUtf8(_)));
}
#[test]
fn truncated_string_fails_to_decode() {
let error = decode_payload::<String>(&[5, b'f']).unwrap_err();
assert!(matches!(error, DecodeError::UnexpectedEof));
}
#[test]
fn uuid_is_sixteen_raw_bytes() {
let uuid = Uuid::from_bytes([7; 16]);
assert_eq!(encode_payload(&uuid).unwrap(), [7; 16]);
assert_eq!(decode_payload::<Uuid>(&[7; 16]).unwrap(), uuid);
let error = decode_payload::<Uuid>(&[0; 15]).unwrap_err();
assert!(matches!(error, DecodeError::UnexpectedEof));
}
#[test]
fn ip_addr_is_version_tagged() {
let v4: IpAddr = Ipv4Addr::new(127, 0, 0, 1).into();
assert_eq!(encode_payload(&v4).unwrap(), [4, 127, 0, 0, 1]);
assert_eq!(decode_payload::<IpAddr>(&[4, 127, 0, 0, 1]).unwrap(), v4);
let v6: IpAddr = Ipv6Addr::LOCALHOST.into();
let mut expected = vec![6];
expected.extend_from_slice(&Ipv6Addr::LOCALHOST.octets());
assert_eq!(encode_payload(&v6).unwrap(), expected);
assert_eq!(decode_payload::<IpAddr>(&expected).unwrap(), v6);
}
#[test]
fn invalid_ip_version_fails_to_decode() {
let error = decode_payload::<IpAddr>(&[0xff, 1, 2, 3, 4]).unwrap_err();
assert!(matches!(error, DecodeError::InvalidIpVersion(0xff)));
}
}