bloop-protocol 1.1.0

Core implementation of the Bloop wire protocol
//! Codec implementations for the spec's primitive data types.

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)));
    }
}