use std::{cmp::Ordering, convert::TryFrom, fmt};
use bytes::BufMut;
use nom::{
IResult,
error::{ErrorKind, make_error},
number::streaming::be_u8,
};
#[derive(Default, Debug, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct VarInt(u64);
pub const VARINT_MAX: u64 = 0x3fff_ffff_ffff_ffff;
#[allow(dead_code)]
pub enum EncodeBytes {
One = 1,
Two = 2,
Four = 4,
Eight = 8,
}
impl VarInt {
pub const MAX: Self = Self(VARINT_MAX);
pub const MAX_SIZE: usize = 8;
pub const fn from_u32(x: u32) -> Self {
Self(x as u64)
}
pub const fn from_u64(value: u64) -> Result<Self, err::Overflow> {
if value <= VARINT_MAX {
Ok(Self(value))
} else {
Err(err::Overflow { value: value as _ })
}
}
pub unsafe fn from_u64_unchecked(x: u64) -> Self {
Self(x)
}
pub fn from_u128(value: u128) -> Result<Self, err::Overflow> {
if value <= VARINT_MAX as u128 {
Ok(Self(value as _))
} else {
Err(err::Overflow { value })
}
}
pub const fn into_inner(self) -> u64 {
self.0
}
pub fn encoding_size(self) -> usize {
let x = self.0;
if x < (1 << 6) {
1
} else if x < (1 << 14) {
2
} else if x < (1 << 30) {
4
} else if x < (1 << 62) {
8
} else {
unreachable!("malformed VarInt");
}
}
}
impl From<VarInt> for u64 {
fn from(x: VarInt) -> Self {
x.0
}
}
impl From<u8> for VarInt {
fn from(x: u8) -> Self {
Self(x.into())
}
}
impl From<u16> for VarInt {
fn from(x: u16) -> Self {
Self(x.into())
}
}
impl From<u32> for VarInt {
fn from(x: u32) -> Self {
Self(x.into())
}
}
impl TryFrom<u128> for VarInt {
type Error = err::Overflow;
fn try_from(x: u128) -> Result<Self, Self::Error> {
Self::from_u128(x)
}
}
impl TryFrom<u64> for VarInt {
type Error = err::Overflow;
fn try_from(x: u64) -> Result<Self, Self::Error> {
Self::from_u64(x)
}
}
impl TryFrom<usize> for VarInt {
type Error = err::Overflow;
fn try_from(x: usize) -> Result<Self, Self::Error> {
Self::try_from(x as u64)
}
}
impl PartialEq<u64> for VarInt {
fn eq(&self, other: &u64) -> bool {
self.0.eq(other)
}
}
impl PartialOrd<u64> for VarInt {
fn partial_cmp(&self, other: &u64) -> Option<Ordering> {
self.0.partial_cmp(other)
}
}
impl fmt::Display for VarInt {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
impl fmt::LowerHex for VarInt {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
impl fmt::UpperHex for VarInt {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
pub mod err {
use std::fmt;
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub struct Overflow {
pub(super) value: u128,
}
impl fmt::Display for Overflow {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Value({}) too large for varint encoding", self.value)
}
}
impl std::error::Error for Overflow {}
}
pub fn be_varint(input: &[u8]) -> IResult<&[u8], VarInt> {
let (remain, first_byte) = be_u8(input)?;
let len = 2usize.pow((first_byte >> 6) as u32);
if remain.len() + 1 < len {
return Err(nom::Err::Incomplete(nom::Needed::new(
len - (remain.len() + 1),
)));
}
let mut buf = [0u8; 8];
buf[0] = first_byte & 0b0011_1111;
buf[1..len].copy_from_slice(&remain[..len - 1]);
let value = u64::from_be_bytes(buf) >> (8 * (8 - len));
Ok((&remain[len - 1..], VarInt(value)))
}
pub trait WriteVarInt {
fn put_varint(&mut self, value: VarInt);
}
impl<T: BufMut> WriteVarInt for T {
fn put_varint(&mut self, VarInt(x): VarInt) {
if x < 1u64 << 6 {
self.put_u8(x as u8);
} else if x < 1u64 << 14 {
self.put_u16((0b01 << 14) | x as u16);
} else if x < 1u64 << 30 {
self.put_u32((0b10 << 30) | x as u32);
} else if x < 1u64 << 62 {
self.put_u64((0b11 << 62) | x);
} else {
unreachable!("malformed VarInt")
};
}
}
#[allow(dead_code)]
pub fn varint_from_u64(value: u64) -> IResult<&'static [u8], VarInt> {
match VarInt::from_u64(value) {
Ok(v) => Ok((&[] as &[u8], v)),
Err(_e) => Err(nom::Err::Error(make_error(
&[] as &[u8],
ErrorKind::TooLarge,
))),
}
}
#[cfg(test)]
mod tests {
use bytes::BytesMut;
use super::*;
fn roundtrip(v: u64) {
let v = VarInt::from_u64(v).unwrap();
let mut buf = BytesMut::new();
buf.put_varint(v);
assert_eq!(buf.len(), v.encoding_size());
let (remain, decoded) = be_varint(&buf).unwrap();
assert!(remain.is_empty());
assert_eq!(decoded, v);
}
#[test]
fn quic_varint_roundtrip() {
for v in [
0u64,
1,
63,
64,
16383,
16384,
(1 << 30) - 1,
1 << 30,
(1 << 62) - 1,
] {
roundtrip(v);
}
}
#[test]
fn quic_varint_rejects_incomplete() {
let v = VarInt::from_u64(64).unwrap();
let mut buf = BytesMut::new();
buf.put_varint(v);
let truncated = &buf[..1];
match be_varint(truncated) {
Err(nom::Err::Incomplete(_)) => {}
other => panic!("expected Incomplete, got {other:?}"),
}
}
#[test]
fn quic_varint_overflow() {
assert!(VarInt::from_u64(VARINT_MAX + 1).is_err());
}
}