use crate::coding::{Decode, DecodeError, Encode, EncodeError, KeyValuePair};
use bytes::Buf;
use std::fmt;
const MIN_EXTENSION_KVP_WIRE_LEN: usize = 2;
#[derive(Default, Clone, Eq, PartialEq)]
pub struct ExtensionHeaders(pub Vec<KeyValuePair>);
impl ExtensionHeaders {
pub fn new() -> Self {
Self::default()
}
pub fn set(&mut self, kvp: KeyValuePair) {
if let Some(existing) = self.0.iter_mut().find(|k| k.key == kvp.key) {
*existing = kvp;
} else {
self.0.push(kvp);
}
}
pub fn set_intvalue(&mut self, key: u64, value: u64) {
self.set(KeyValuePair::new_int(key, value));
}
pub fn set_bytesvalue(&mut self, key: u64, value: Vec<u8>) {
self.set(KeyValuePair::new_bytes(key, value));
}
pub fn has(&self, key: u64) -> bool {
self.0.iter().any(|k| k.key == key)
}
pub fn get(&self, key: u64) -> Option<&KeyValuePair> {
self.0.iter().find(|k| k.key == key)
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl Decode for ExtensionHeaders {
fn decode<R: bytes::Buf>(r: &mut R) -> Result<Self, DecodeError> {
let length = usize::decode(r)?;
Self::decode_remaining(r, length)?;
if length == 0 {
return Ok(ExtensionHeaders::new());
}
let mut kvps_bytes = r.copy_to_bytes(length);
let mut kvps = Vec::with_capacity(length / MIN_EXTENSION_KVP_WIRE_LEN);
let mut prev = 0u64;
while kvps_bytes.has_remaining() {
let (pair, new_prev) = KeyValuePair::decode_with_prev(&mut kvps_bytes, prev)?;
prev = new_prev;
kvps.push(pair);
}
Ok(ExtensionHeaders(kvps))
}
}
impl Encode for ExtensionHeaders {
fn encode<W: bytes::BufMut>(&self, w: &mut W) -> Result<(), EncodeError> {
if self.0.is_empty() {
0usize.encode(w)?;
return Ok(());
}
let mut tmp = bytes::BytesMut::new();
if let [kvp] = self.0.as_slice() {
kvp.encode_with_prev(&mut tmp, 0)?;
} else {
let mut sorted: Vec<&KeyValuePair> = self.0.iter().collect();
sorted.sort_by_key(|k| k.key);
let mut prev = 0u64;
for kvp in &sorted {
prev = kvp.encode_with_prev(&mut tmp, prev)?;
}
}
tmp.len().encode(w)?;
w.put_slice(&tmp);
Ok(())
}
}
impl fmt::Debug for ExtensionHeaders {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{{ ")?;
for (i, kv) in self.0.iter().enumerate() {
if i > 0 {
write!(f, ", ")?;
}
write!(f, "{:?}", kv)?;
}
write!(f, " }}")
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::BytesMut;
#[test]
fn single_bytes_pair_roundtrip() {
let mut buf = BytesMut::new();
let mut ext = ExtensionHeaders::new();
ext.set_bytesvalue(1, vec![0x01, 0x02, 0x03, 0x04, 0x05]);
ext.encode(&mut buf).unwrap();
assert_eq!(
buf.to_vec(),
vec![
0x07, 0x01, 0x05, 0x01, 0x02, 0x03, 0x04, 0x05, ]
);
let decoded = ExtensionHeaders::decode(&mut buf).unwrap();
assert_eq!(decoded, ext);
}
#[test]
fn multi_pair_delta_encoding() {
let mut ext = ExtensionHeaders::new();
ext.set_intvalue(0, 0);
ext.set_intvalue(100, 100);
ext.set_bytesvalue(1, vec![0x01, 0x02, 0x03, 0x04, 0x05]);
let mut buf = BytesMut::new();
ext.encode(&mut buf).unwrap();
let buf_vec = buf.to_vec();
assert_eq!(buf_vec[0], 13); assert_eq!(buf_vec.len(), 14);
let decoded = ExtensionHeaders::decode(&mut buf).unwrap();
assert_eq!(decoded.0.len(), 3);
assert_eq!(
decoded.get(0).unwrap().value,
crate::coding::Value::IntValue(0)
);
assert_eq!(
decoded.get(100).unwrap().value,
crate::coding::Value::IntValue(100)
);
assert_eq!(
decoded.get(1).unwrap().value,
crate::coding::Value::BytesValue(vec![0x01, 0x02, 0x03, 0x04, 0x05])
);
}
#[test]
fn encode_sorts_before_delta() {
let mut ext = ExtensionHeaders::new();
ext.set_intvalue(100, 99);
ext.set_intvalue(0, 1);
let mut buf = BytesMut::new();
ext.encode(&mut buf).unwrap();
let decoded = ExtensionHeaders::decode(&mut buf).unwrap();
assert_eq!(
decoded.get(0).unwrap().value,
crate::coding::Value::IntValue(1)
);
assert_eq!(
decoded.get(100).unwrap().value,
crate::coding::Value::IntValue(99)
);
}
#[test]
fn empty_roundtrip() {
let ext = ExtensionHeaders::new();
let mut buf = BytesMut::new();
ext.encode(&mut buf).unwrap();
assert_eq!(buf.to_vec(), vec![0x00]); let decoded = ExtensionHeaders::decode(&mut buf).unwrap();
assert_eq!(decoded, ext);
}
}