use crate::error::{Error, ErrorCode, Result, fmt};
pub const MAX_VARINT_LEN_U64: usize = 10;
pub fn encode_u64(mut value: u64, out: &mut Vec<u8>) -> usize {
let start = out.len();
while value & !0x7F != 0 {
out.push(((value & 0x7F) as u8) | 0x80);
value >>= 7;
}
out.push(value as u8);
out.len() - start
}
pub fn decode_u64(bytes: &[u8]) -> Result<(u64, usize)> {
let mut result: u64 = 0;
let mut shift: u32 = 0;
for (i, &b) in bytes.iter().enumerate() {
if shift >= 64 {
return Err(fmt!(
ProtocolError,
"varint exceeds 64-bit range at byte {}",
i
));
}
let chunk = (b & 0x7F) as u64;
if shift == 63 && (chunk & !0x01) != 0 {
return Err(fmt!(
ProtocolError,
"varint exceeds 64-bit range at byte {}",
i
));
}
result |= chunk << shift;
if b & 0x80 == 0 {
return Ok((result, i + 1));
}
shift += 7;
}
Err(fmt!(
ProtocolError,
"truncated varint: {} bytes without terminator",
bytes.len()
))
}
pub fn decode_usize(bytes: &[u8]) -> Result<(usize, usize)> {
let (v, n) = decode_u64(bytes)?;
let v_us = usize::try_from(v).map_err(|_| {
Error::new(
ErrorCode::ProtocolError,
format!("varint value {} does not fit in usize", v),
)
})?;
Ok((v_us, n))
}
#[cfg(test)]
mod tests {
use super::*;
fn roundtrip(value: u64, expected_len: usize) {
let mut buf = Vec::new();
let n = encode_u64(value, &mut buf);
assert_eq!(n, expected_len, "encoded length for {}", value);
assert_eq!(buf.len(), expected_len);
let (decoded, consumed) = decode_u64(&buf).expect("decode");
assert_eq!(decoded, value);
assert_eq!(consumed, expected_len);
}
#[test]
fn boundaries() {
roundtrip(0, 1);
roundtrip(1, 1);
roundtrip(0x7F, 1);
roundtrip(0x80, 2);
roundtrip(0x3FFF, 2);
roundtrip(0x4000, 3);
roundtrip(u32::MAX as u64, 5);
roundtrip(u64::MAX, 10);
}
#[test]
fn reference_vector_300() {
let mut buf = Vec::new();
encode_u64(300, &mut buf);
assert_eq!(buf, vec![0xAC, 0x02]);
}
#[test]
fn truncated_is_error() {
let bytes = [0x80u8];
let err = decode_u64(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn overlong_is_error() {
let bytes = [0x80u8; 11];
let err = decode_u64(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn tenth_byte_with_high_bits_is_error() {
let bytes = [0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x02];
let err = decode_u64(&bytes).unwrap_err();
assert_eq!(err.code(), ErrorCode::ProtocolError);
}
#[test]
fn decode_consumes_only_one_value() {
let mut buf = Vec::new();
encode_u64(300, &mut buf);
encode_u64(7, &mut buf);
let (v1, n1) = decode_u64(&buf).unwrap();
assert_eq!(v1, 300);
let (v2, n2) = decode_u64(&buf[n1..]).unwrap();
assert_eq!(v2, 7);
assert_eq!(n1 + n2, buf.len());
}
#[test]
fn decode_usize_succeeds_for_small_values() {
let mut buf = Vec::new();
encode_u64(42, &mut buf);
let (v, _) = decode_usize(&buf).unwrap();
assert_eq!(v, 42);
}
}