use crate::coding::{Decode, DecodeError, Encode, EncodeError};
use std::fmt;
const MAX_BYTES_VALUE_LEN: usize = u16::MAX as usize;
const MIN_KVP_WIRE_LEN: usize = 2;
#[derive(Clone, Eq, PartialEq)]
pub enum Value {
IntValue(u64),
BytesValue(Vec<u8>),
}
impl fmt::Debug for Value {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Value::IntValue(v) => write!(f, "{}", v),
Value::BytesValue(bytes) => {
let preview: Vec<String> = bytes
.iter()
.take(16)
.map(|b| format!("{:02X}", b))
.collect();
write!(f, "[{}]", preview.join(" "))
}
}
}
}
#[derive(Clone, Eq, PartialEq)]
pub struct KeyValuePair {
pub key: u64,
pub value: Value,
}
impl KeyValuePair {
pub fn new(key: u64, value: Value) -> Self {
Self { key, value }
}
pub fn new_int(key: u64, value: u64) -> Self {
Self {
key,
value: Value::IntValue(value),
}
}
pub fn new_bytes(key: u64, value: Vec<u8>) -> Self {
Self {
key,
value: Value::BytesValue(value),
}
}
pub(crate) fn decode_with_prev<R: bytes::Buf>(
r: &mut R,
prev: u64,
) -> Result<(Self, u64), DecodeError> {
let delta = u64::decode(r)?;
let abs_type = prev
.checked_add(delta)
.ok_or(DecodeError::KvpTypeOverflow)?;
let pair = if abs_type % 2 == 0 {
let value = u64::decode(r)?;
KeyValuePair::new_int(abs_type, value)
} else {
let length = usize::decode(r)?;
if length > MAX_BYTES_VALUE_LEN {
return Err(DecodeError::KeyValuePairLengthExceeded());
}
<u64 as Decode>::decode_remaining(r, length)?;
let mut buf = vec![0u8; length];
r.copy_to_slice(&mut buf);
KeyValuePair::new_bytes(abs_type, buf)
};
Ok((pair, abs_type))
}
pub(crate) fn encode_with_prev<W: bytes::BufMut>(
&self,
w: &mut W,
prev: u64,
) -> Result<u64, EncodeError> {
match &self.value {
Value::IntValue(_) if !self.key.is_multiple_of(2) => {
return Err(EncodeError::InvalidValue);
}
Value::BytesValue(_) if self.key.is_multiple_of(2) => {
return Err(EncodeError::InvalidValue);
}
_ => {}
}
let delta = self.key.checked_sub(prev).ok_or(EncodeError::KvpKeyOrder)?;
delta.encode(w)?;
match &self.value {
Value::IntValue(v) => {
(*v).encode(w)?;
}
Value::BytesValue(v) => {
if v.len() > MAX_BYTES_VALUE_LEN {
return Err(EncodeError::KeyValuePairLengthExceeded);
}
v.len().encode(w)?;
<u64 as Encode>::encode_remaining(w, v.len())?;
w.put_slice(v);
}
}
Ok(self.key)
}
}
impl fmt::Debug for KeyValuePair {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{{{}: {:?}}}", self.key, self.value)
}
}
#[derive(Default, Clone, Eq, PartialEq)]
pub struct KeyValuePairs(pub Vec<KeyValuePair>);
impl KeyValuePairs {
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 has_duplicate_keys(&self) -> bool {
let mut seen = std::collections::HashSet::new();
self.0.iter().any(|k| !seen.insert(k.key))
}
}
impl Decode for KeyValuePairs {
fn decode<R: bytes::Buf>(r: &mut R) -> Result<Self, DecodeError> {
let count = u64::decode(r)?;
let count_capacity = usize::try_from(count).unwrap_or(usize::MAX);
let payload_capacity = r.remaining() / MIN_KVP_WIRE_LEN;
let mut kvps = Vec::with_capacity(count_capacity.min(payload_capacity));
let mut prev = 0u64;
for _ in 0..count {
let (pair, new_prev) = KeyValuePair::decode_with_prev(r, prev)?;
prev = new_prev;
kvps.push(pair);
}
Ok(KeyValuePairs(kvps))
}
}
impl Encode for KeyValuePairs {
fn encode<W: bytes::BufMut>(&self, w: &mut W) -> Result<(), EncodeError> {
let mut sorted: Vec<&KeyValuePair> = self.0.iter().collect();
sorted.sort_by_key(|k| k.key);
sorted.len().encode(w)?;
let mut prev = 0u64;
for kvp in &sorted {
prev = kvp.encode_with_prev(w, prev)?;
}
Ok(())
}
}
impl fmt::Debug for KeyValuePairs {
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;
fn round_trip_pair(pairs: &[(u64, Value)]) -> Vec<u8> {
let mut buf = BytesMut::new();
let mut prev = 0u64;
for (key, value) in pairs {
let delta = key - prev;
delta.encode(&mut buf).unwrap();
match value {
Value::IntValue(v) => v.encode(&mut buf).unwrap(),
Value::BytesValue(b) => {
b.len().encode(&mut buf).unwrap();
buf.extend_from_slice(b);
}
}
prev = *key;
}
buf.to_vec()
}
#[test]
fn single_int_pair_roundtrip() {
let mut buf = BytesMut::new();
let kvps = KeyValuePairs(vec![KeyValuePair::new_int(0, 42)]);
kvps.encode(&mut buf).unwrap();
assert_eq!(buf.to_vec(), vec![0x01, 0x00, 0x2a]);
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded, kvps);
}
#[test]
fn single_bytes_pair_roundtrip() {
let mut buf = BytesMut::new();
let kvps = KeyValuePairs(vec![KeyValuePair::new_bytes(1, vec![0xAB, 0xCD])]);
kvps.encode(&mut buf).unwrap();
assert_eq!(buf.to_vec(), vec![0x01, 0x01, 0x02, 0xAB, 0xCD]);
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded, kvps);
}
#[test]
fn delta_encoding_multiple_pairs() {
let mut buf = BytesMut::new();
let mut kvps = KeyValuePairs::new();
kvps.set_intvalue(0, 1);
kvps.set_intvalue(2, 2);
kvps.set_intvalue(100, 3);
kvps.encode(&mut buf).unwrap();
let expected_wire = round_trip_pair(&[
(0, Value::IntValue(1)),
(2, Value::IntValue(2)),
(100, Value::IntValue(3)),
]);
assert_eq!(buf[1..], expected_wire[..]);
assert_eq!(buf[0], 0x03);
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded.0.len(), 3);
assert_eq!(decoded.get(0).unwrap().value, Value::IntValue(1));
assert_eq!(decoded.get(2).unwrap().value, Value::IntValue(2));
assert_eq!(decoded.get(100).unwrap().value, Value::IntValue(3));
}
#[test]
fn encode_sorts_before_delta() {
let mut kvps = KeyValuePairs::new();
kvps.set_intvalue(100, 99);
kvps.set_intvalue(0, 1);
kvps.set_intvalue(2, 2);
let mut buf = BytesMut::new();
kvps.encode(&mut buf).unwrap();
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded.0.len(), 3);
assert_eq!(decoded.get(0).unwrap().value, Value::IntValue(1));
assert_eq!(decoded.get(2).unwrap().value, Value::IntValue(2));
assert_eq!(decoded.get(100).unwrap().value, Value::IntValue(99));
}
#[test]
fn encode_rejects_odd_key_with_int_value() {
let kvp = KeyValuePair::new(1, Value::IntValue(0)); let mut buf = BytesMut::new();
let kvps = KeyValuePairs(vec![kvp]);
assert!(matches!(
kvps.encode(&mut buf).unwrap_err(),
EncodeError::InvalidValue
));
}
#[test]
fn encode_rejects_even_key_with_bytes_value() {
let kvp = KeyValuePair::new(0, Value::BytesValue(vec![0x01])); let kvps = KeyValuePairs(vec![kvp]);
let mut buf = BytesMut::new();
assert!(matches!(
kvps.encode(&mut buf).unwrap_err(),
EncodeError::InvalidValue
));
}
#[test]
fn decode_detects_delta_overflow() {
let max_delta: u64 = (1u64 << 62) - 1;
let mut buf = BytesMut::new();
(5u64).encode(&mut buf).unwrap();
max_delta.encode(&mut buf).unwrap();
(0usize).encode(&mut buf).unwrap();
max_delta.encode(&mut buf).unwrap();
(0u64).encode(&mut buf).unwrap();
max_delta.encode(&mut buf).unwrap();
(0usize).encode(&mut buf).unwrap();
max_delta.encode(&mut buf).unwrap();
(0u64).encode(&mut buf).unwrap();
max_delta.encode(&mut buf).unwrap();
let err = KeyValuePairs::decode(&mut buf).unwrap_err();
assert!(
matches!(err, DecodeError::KvpTypeOverflow),
"expected KvpTypeOverflow, got {:?}",
err
);
}
#[test]
fn decode_rejects_bytes_value_too_long() {
let mut buf = BytesMut::new();
(1u64).encode(&mut buf).unwrap(); (1u64).encode(&mut buf).unwrap(); let too_long = MAX_BYTES_VALUE_LEN + 1;
too_long.encode(&mut buf).unwrap();
let err = KeyValuePairs::decode(&mut buf).unwrap_err();
assert!(
matches!(err, DecodeError::KeyValuePairLengthExceeded()),
"expected KeyValuePairLengthExceeded, got {:?}",
err
);
}
#[test]
fn decode_large_count_does_not_allocate_count_capacity() {
let mut buf = BytesMut::new();
((1u64 << 62) - 1).encode(&mut buf).unwrap();
let err = KeyValuePairs::decode(&mut buf).unwrap_err();
assert!(matches!(err, DecodeError::More(_)));
}
#[test]
fn has_duplicate_keys_detects_duplicates() {
let mut kvps = KeyValuePairs::new();
kvps.0.push(KeyValuePair::new_int(0, 1));
kvps.0.push(KeyValuePair::new_int(0, 2)); assert!(kvps.has_duplicate_keys());
}
#[test]
fn has_duplicate_keys_no_false_positive() {
let mut kvps = KeyValuePairs::new();
kvps.set_intvalue(0, 1);
kvps.set_intvalue(2, 2);
assert!(!kvps.has_duplicate_keys());
}
#[test]
fn empty_kvps_roundtrip() {
let kvps = KeyValuePairs::new();
let mut buf = BytesMut::new();
kvps.encode(&mut buf).unwrap();
assert_eq!(buf.to_vec(), vec![0x00]); let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded, kvps);
}
#[test]
fn existing_single_bytes_compat() {
let mut buf = BytesMut::new();
let mut kvps = KeyValuePairs::new();
kvps.set_bytesvalue(1, vec![0x01, 0x02, 0x03, 0x04, 0x05]);
kvps.encode(&mut buf).unwrap();
assert_eq!(
buf.to_vec(),
vec![
0x01, 0x01, 0x05, 0x01, 0x02, 0x03, 0x04, 0x05, ]
);
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded, kvps);
}
#[test]
fn existing_multi_compat() {
let mut buf = BytesMut::new();
let mut kvps = KeyValuePairs::new();
kvps.set_intvalue(0, 0);
kvps.set_intvalue(100, 100);
kvps.set_bytesvalue(1, vec![0x01, 0x02, 0x03, 0x04, 0x05]);
kvps.encode(&mut buf).unwrap();
let decoded = KeyValuePairs::decode(&mut buf).unwrap();
assert_eq!(decoded.0.len(), 3);
assert_eq!(decoded.get(0).unwrap().value, Value::IntValue(0));
assert_eq!(decoded.get(100).unwrap().value, Value::IntValue(100));
assert_eq!(
decoded.get(1).unwrap().value,
Value::BytesValue(vec![0x01, 0x02, 0x03, 0x04, 0x05])
);
}
}