use std::io;
use std::io::Write;
use nom::{IResult, be_u8};
pub const NUM_EXTENSION_BYTES: usize = 8;
pub enum Extension {
ExtensionProtocol = 43
}
#[derive(Copy, Clone, Eq, Hash, PartialEq, Debug)]
pub struct Extensions {
bytes: [u8; NUM_EXTENSION_BYTES]
}
impl Extensions {
pub fn new() -> Extensions {
Extensions::with_bytes([0u8; NUM_EXTENSION_BYTES])
}
pub fn from_bytes(bytes: &[u8]) -> IResult<&[u8], Extensions> {
parse_extension_bits(bytes)
}
pub fn add(&mut self, extension: Extension) {
let active_bit = extension as usize;
let byte_index = active_bit / 8;
let bit_index = active_bit % 8;
self.bytes[byte_index] |= 0x80 >> bit_index;
}
pub fn remove(&mut self, extension: Extension) {
let active_bit = extension as usize;
let byte_index = active_bit / 8;
let bit_index = active_bit % 8;
self.bytes[byte_index] &= !(0x80 >> bit_index);
}
pub fn contains(&self, extension: Extension) -> bool {
let active_bit = extension as usize;
let byte_index = active_bit / 8;
let bit_index = active_bit % 8;
self.bytes[byte_index] & (0x80 >> bit_index) != 0
}
pub fn write_bytes<W>(&self, mut writer: W) -> io::Result<()>
where W: Write {
writer.write_all(&self.bytes[..])
}
pub fn union(&self, ext: &Extensions) -> Extensions {
let mut result_ext = Extensions::new();
for index in 0..NUM_EXTENSION_BYTES {
result_ext.bytes[index] = self.bytes[index] & ext.bytes[index];
}
result_ext
}
fn with_bytes(bytes: [u8; NUM_EXTENSION_BYTES]) -> Extensions {
Extensions{ bytes: bytes }
}
}
impl From<[u8; NUM_EXTENSION_BYTES]> for Extensions {
fn from(bytes: [u8; NUM_EXTENSION_BYTES]) -> Extensions {
Extensions{ bytes: bytes }
}
}
fn parse_extension_bits(bytes: &[u8]) -> IResult<&[u8], Extensions> {
do_parse!(bytes,
bytes: count_fixed!(u8, be_u8, NUM_EXTENSION_BYTES) >>
(Extensions::with_bytes(bytes))
)
}
#[cfg(test)]
mod tests {
use super::{Extensions, Extension};
#[test]
fn positive_add_extension_protocol() {
let mut extensions = Extensions::new();
extensions.add(Extension::ExtensionProtocol);
let expected_extensions: Extensions = [0, 0, 0, 0, 0, 0x10, 0, 0].into();
assert_eq!(expected_extensions, extensions);
assert!(extensions.contains(Extension::ExtensionProtocol));
}
#[test]
fn positive_remove_extension_protocol() {
let mut extensions = Extensions::new();
extensions.add(Extension::ExtensionProtocol);
extensions.remove(Extension::ExtensionProtocol);
let expected_extensions: Extensions = [0, 0, 0, 0, 0, 0, 0, 0].into();
assert_eq!(expected_extensions, extensions);
assert!(!extensions.contains(Extension::ExtensionProtocol));
}
}