use bytes::Buf as _;
use crate::coding::{Decode, DecodeError, Encode, EncodeError};
use super::{Location, Param, Version};
#[derive(Default, Clone, Copy, Debug, PartialEq, Eq)]
pub enum Filter {
Unfiltered,
#[default]
NextObject,
Relative(u64),
Absolute {
start: Location,
end: Option<EndLocation>,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct EndLocation {
pub group: u64,
pub object: Option<u64>,
}
mod tag {
pub const NEXT_GROUP: u64 = 0x1;
pub const LARGEST_OBJECT: u64 = 0x2;
pub const ABSOLUTE_START: u64 = 0x3;
pub const ABSOLUTE_RANGE: u64 = 0x4;
}
impl Filter {
pub(crate) fn is_draft20(version: Version) -> bool {
!matches!(
version,
Version::Draft14
| Version::Draft15
| Version::Draft16
| Version::Draft17
| Version::Draft18
| Version::Draft19
)
}
fn end_delta(start: u64, end: u64) -> Result<u64, EncodeError> {
end.checked_sub(start).ok_or(EncodeError::InvalidState)
}
fn encode_fields<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match *self {
Self::Unfiltered => {}
Self::NextObject => {
0u64.encode(w, version)?;
0u64.encode(w, version)?;
}
Self::Relative(groups) => groups.encode(w, version)?,
Self::Absolute {
start: Location { group: 0, object: 0 },
end: None,
} => {}
Self::Absolute { start, end } => {
start.group.encode(w, version)?;
start.object.encode(w, version)?;
if let Some(end) = end {
Self::end_delta(start.group, end.group)?.encode(w, version)?;
if let Some(object) = end.object {
object.encode(w, version)?;
}
}
}
}
Ok(())
}
fn decode_fields<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let mut fields = Vec::with_capacity(4);
while r.has_remaining() {
if fields.len() == 4 {
return Err(DecodeError::TrailingBytes);
}
fields.push(u64::decode(r, version)?);
}
Ok(match fields[..] {
[] => Self::Unfiltered,
[groups] => Self::Relative(groups),
[0, 0] => Self::NextObject,
[group, object] => Self::Absolute {
start: Location { group, object },
end: None,
},
[group, object, delta] => Self::Absolute {
start: Location { group, object },
end: Some(EndLocation {
group: group.checked_add(delta).ok_or(DecodeError::BoundsExceeded)?,
object: None,
}),
},
[group, object, delta, end_object] => Self::Absolute {
start: Location { group, object },
end: Some(EndLocation {
group: group.checked_add(delta).ok_or(DecodeError::BoundsExceeded)?,
object: Some(end_object),
}),
},
_ => unreachable!("capped at 4 fields above"),
})
}
fn encode_tag<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match *self {
Self::Unfiltered => {
tag::ABSOLUTE_START.encode(w, version)?;
Location::default().encode(w, version)?;
}
Self::NextObject => tag::LARGEST_OBJECT.encode(w, version)?,
Self::Relative(0) => tag::NEXT_GROUP.encode(w, version)?,
Self::Relative(_) => return Err(EncodeError::Unsupported),
Self::Absolute { start, end: None } => {
tag::ABSOLUTE_START.encode(w, version)?;
start.encode(w, version)?;
}
Self::Absolute {
end: Some(EndLocation { object: Some(_), .. }),
..
} => return Err(EncodeError::Unsupported),
Self::Absolute { start, end: Some(end) } => {
tag::ABSOLUTE_RANGE.encode(w, version)?;
start.encode(w, version)?;
Self::end_delta(start.group, end.group)?.encode(w, version)?;
}
}
Ok(())
}
fn decode_tag<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
Ok(match u64::decode(r, version)? {
tag::NEXT_GROUP => Self::Relative(0),
tag::LARGEST_OBJECT => Self::NextObject,
tag::ABSOLUTE_START => Self::Absolute {
start: Location::decode(r, version)?,
end: None,
},
tag::ABSOLUTE_RANGE => {
let start = Location::decode(r, version)?;
let delta = u64::decode(r, version)?;
Self::Absolute {
start,
end: Some(EndLocation {
group: start.group.checked_add(delta).ok_or(DecodeError::BoundsExceeded)?,
object: None,
}),
}
}
_ => return Err(DecodeError::InvalidValue),
})
}
}
impl Encode<Version> for Filter {
fn encode<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
self.encode_tag(w, version)
}
}
impl Decode<Version> for Filter {
fn decode<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
Self::decode_tag(r, version)
}
}
impl Param for Filter {
fn param_present(&self) -> bool {
!matches!(self, Self::Unfiltered)
}
fn param_encode<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
let mut buf = Vec::new();
let sv = match version {
Version::Draft14 | Version::Draft15 | Version::Draft16 => Version::Draft15,
_ => version,
};
if Self::is_draft20(version) {
self.encode_fields(&mut buf, sv)?;
} else {
self.encode_tag(&mut buf, sv)?;
}
buf.encode(w, version)
}
fn param_decode<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let data = Vec::<u8>::decode(r, version)?;
let mut buf = bytes::Bytes::from(data);
let sv = match version {
Version::Draft14 | Version::Draft15 | Version::Draft16 => Version::Draft15,
_ => version,
};
if Self::is_draft20(version) {
return Self::decode_fields(&mut buf, sv);
}
let filter = Self::decode_tag(&mut buf, sv)?;
if buf.has_remaining() {
return Err(DecodeError::TrailingBytes);
}
Ok(filter)
}
}
#[cfg(test)]
mod tests {
use super::*;
const OLD: Version = Version::Draft19;
const NEW: Version = Version::Draft20;
fn round_trip(filter: Filter, version: Version) -> Filter {
let mut buf = Vec::new();
filter.param_encode(&mut buf, version).expect("encode");
let mut bytes = bytes::Bytes::from(buf);
let decoded = Filter::param_decode(&mut bytes, version).expect("decode");
assert!(!bytes.has_remaining(), "parameter left trailing bytes");
decoded
}
fn value(filter: Filter, version: Version) -> Vec<u8> {
let mut buf = Vec::new();
filter.param_encode(&mut buf, version).expect("encode");
let mut bytes = bytes::Bytes::from(buf);
Vec::<u8>::decode(&mut bytes, version).expect("length prefix")
}
#[test]
fn draft20_field_counts() {
assert_eq!(value(Filter::NextObject, NEW), vec![0x00, 0x00]);
assert_eq!(value(Filter::Relative(0), NEW), vec![0x00]);
assert_eq!(value(Filter::Relative(1), NEW), vec![0x01]);
assert_eq!(
value(
Filter::Absolute {
start: Location { group: 7, object: 3 },
end: None
},
NEW
),
vec![0x07, 0x03]
);
assert_eq!(
value(
Filter::Absolute {
start: Location { group: 7, object: 3 },
end: Some(EndLocation { group: 9, object: None })
},
NEW
),
vec![0x07, 0x03, 0x02]
);
}
#[test]
fn draft20_uses_leading_ones_varints_above_63() {
assert_eq!(value(Filter::Relative(128), NEW), vec![0x80, 0x80]);
assert_eq!(round_trip(Filter::Relative(128), NEW), Filter::Relative(128));
let wide = Filter::Absolute {
start: Location { group: 300, object: 64 },
end: None,
};
assert_eq!(round_trip(wide, NEW), wide);
}
#[test]
fn draft20_names_the_current_group() {
assert_eq!(value(Filter::Relative(1), NEW), vec![0x01]);
assert_eq!(round_trip(Filter::Relative(1), NEW), Filter::Relative(1));
}
#[test]
fn draft20_round_trips() {
for filter in [
Filter::Unfiltered,
Filter::NextObject,
Filter::Relative(0),
Filter::Relative(5),
Filter::Absolute {
start: Location { group: 12, object: 0 },
end: None,
},
Filter::Absolute {
start: Location { group: 12, object: 4 },
end: Some(EndLocation {
group: 20,
object: None,
}),
},
] {
assert_eq!(round_trip(filter, NEW), filter, "{filter:?}");
}
}
#[test]
fn draft20_absolute_origin_is_unfiltered() {
let origin = Filter::Absolute {
start: Location::default(),
end: None,
};
assert!(value(origin, NEW).is_empty());
assert_eq!(round_trip(origin, NEW), Filter::Unfiltered);
assert_ne!(value(origin, NEW), value(Filter::NextObject, NEW));
}
#[test]
fn draft19_uses_tags() {
assert_eq!(value(Filter::NextObject, OLD), vec![tag::LARGEST_OBJECT as u8]);
assert_eq!(value(Filter::Relative(0), OLD), vec![tag::NEXT_GROUP as u8]);
for filter in [
Filter::NextObject,
Filter::Relative(0),
Filter::Absolute {
start: Location { group: 12, object: 4 },
end: None,
},
Filter::Absolute {
start: Location { group: 12, object: 4 },
end: Some(EndLocation {
group: 20,
object: None,
}),
},
] {
assert_eq!(round_trip(filter, OLD), filter, "{filter:?}");
}
}
#[test]
fn draft20_keeps_the_end_object() {
let bounded = Filter::Absolute {
start: Location { group: 7, object: 3 },
end: Some(EndLocation {
group: 9,
object: Some(4),
}),
};
assert_eq!(value(bounded, NEW), vec![0x07, 0x03, 0x02, 0x04]);
assert_eq!(round_trip(bounded, NEW), bounded);
let whole = Filter::Absolute {
start: Location { group: 7, object: 3 },
end: Some(EndLocation { group: 9, object: None }),
};
assert_eq!(value(whole, NEW), vec![0x07, 0x03, 0x02]);
assert_ne!(round_trip(bounded, NEW), whole);
}
#[test]
fn draft19_cannot_bound_the_end_object() {
let mut buf = Vec::new();
let bounded = Filter::Absolute {
start: Location { group: 7, object: 0 },
end: Some(EndLocation {
group: 9,
object: Some(4),
}),
};
assert!(bounded.param_encode(&mut buf, OLD).is_err());
}
#[test]
fn draft19_cannot_name_a_relative_group() {
let mut buf = Vec::new();
assert!(Filter::Relative(2).param_encode(&mut buf, OLD).is_err());
}
#[test]
fn rejects_a_backwards_range() {
let backwards = Filter::Absolute {
start: Location { group: 9, object: 0 },
end: Some(EndLocation { group: 4, object: None }),
};
for version in [OLD, NEW] {
let mut buf = Vec::new();
assert!(backwards.param_encode(&mut buf, version).is_err(), "{version}");
}
}
#[test]
fn rejects_too_many_fields() {
let mut buf = Vec::new();
vec![0u8; 5].encode(&mut buf, NEW).expect("encode");
let mut bytes = bytes::Bytes::from(buf);
assert!(Filter::param_decode(&mut bytes, NEW).is_err());
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Framing {
Byte,
Varint,
Bytes,
}
#[derive(Default, Clone, Copy, Debug, PartialEq, Eq)]
pub struct Fill {
pub filter: Option<Filter>,
pub range_filters: bool,
}
impl Fill {
const LOCATION_FILTER: u64 = 0x21;
const ALLOWED: &'static [(u64, Framing)] = &[
(0x0A, Framing::Varint), (0x20, Framing::Byte), (Self::LOCATION_FILTER, Framing::Bytes),
(0x22, Framing::Byte), (0x25, Framing::Bytes), (0x26, Framing::Bytes), (0x27, Framing::Bytes), (0x28, Framing::Bytes), ];
}
impl Param for Fill {
fn param_encode<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
let mut buf = Vec::new();
match self.filter {
None => 0u64.encode(&mut buf, version)?,
Some(filter) => {
1u64.encode(&mut buf, version)?;
Self::LOCATION_FILTER.encode(&mut buf, version)?;
filter.param_encode(&mut buf, version)?;
}
}
buf.encode(w, version)
}
fn param_decode<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let data = Vec::<u8>::decode(r, version)?;
let mut buf = bytes::Bytes::from(data);
let count = u64::decode(&mut buf, version)?;
if count > 64 {
return Err(DecodeError::TooMany);
}
let mut filter = None;
let mut range_filters = false;
let mut prev = 0u64;
for i in 0..count {
let delta = u64::decode(&mut buf, version)?;
let key = if i == 0 {
delta
} else {
prev.checked_add(delta).ok_or(DecodeError::BoundsExceeded)?
};
prev = key;
let Some((_, framing)) = Self::ALLOWED.iter().find(|(id, _)| *id == key) else {
return Err(DecodeError::InvalidValue);
};
if key == Self::LOCATION_FILTER {
if filter.is_some() {
return Err(DecodeError::Duplicate);
}
filter = Some(Filter::param_decode(&mut buf, version)?);
continue;
}
range_filters |= (0x25..=0x28).contains(&key);
match framing {
Framing::Bytes => {
Vec::<u8>::decode(&mut buf, version)?;
}
Framing::Varint => {
u64::decode(&mut buf, version)?;
}
Framing::Byte => {
u8::decode(&mut buf, version)?;
}
}
}
if buf.has_remaining() {
return Err(DecodeError::TrailingBytes);
}
Ok(Self { filter, range_filters })
}
}
#[cfg(test)]
mod fill_tests {
use super::*;
const NEW: Version = Version::Draft20;
fn round_trip(fill: Fill) -> Fill {
let mut buf = Vec::new();
fill.param_encode(&mut buf, NEW).expect("encode");
let mut bytes = bytes::Bytes::from(buf);
let decoded = Fill::param_decode(&mut bytes, NEW).expect("decode");
assert!(!bytes.has_remaining());
decoded
}
#[test]
fn round_trips() {
for filter in [
None,
Some(Filter::Unfiltered),
Some(Filter::Relative(1)),
Some(Filter::Relative(3)),
Some(Filter::Absolute {
start: Location { group: 4, object: 0 },
end: Some(EndLocation { group: 9, object: None }),
}),
] {
let fill = Fill {
filter,
range_filters: false,
};
assert_eq!(round_trip(fill).filter, filter, "{filter:?}");
}
}
#[test]
fn current_group_join() {
let fill = Fill {
filter: Some(Filter::Relative(1)),
range_filters: false,
};
let mut buf = Vec::new();
fill.param_encode(&mut buf, NEW).expect("encode");
let mut bytes = bytes::Bytes::from(buf);
let value = Vec::<u8>::decode(&mut bytes, NEW).expect("length prefix");
assert_eq!(value, vec![0x01, 0x21, 0x01, 0x01]);
}
#[test]
fn rejects_a_disallowed_parameter() {
let mut value = Vec::new();
1u64.encode(&mut value, NEW).unwrap();
0x10u64.encode(&mut value, NEW).unwrap(); 0u64.encode(&mut value, NEW).unwrap();
let mut buf = Vec::new();
value.encode(&mut buf, NEW).unwrap();
let mut bytes = bytes::Bytes::from(buf);
assert!(Fill::param_decode(&mut bytes, NEW).is_err());
}
#[test]
fn skips_a_uint8_whose_value_has_a_leading_one() {
let mut value = Vec::new();
2u64.encode(&mut value, NEW).unwrap();
0x20u64.encode(&mut value, NEW).unwrap(); 0x80u8.encode(&mut value, NEW).unwrap(); 1u64.encode(&mut value, NEW).unwrap(); Filter::Relative(1).param_encode(&mut value, NEW).unwrap();
let mut buf = Vec::new();
value.encode(&mut buf, NEW).unwrap();
let mut bytes = bytes::Bytes::from(buf);
let fill = Fill::param_decode(&mut bytes, NEW).expect("decode");
assert_eq!(fill.filter, Some(Filter::Relative(1)));
}
#[test]
fn skips_a_length_prefixed_range_filter() {
let mut value = Vec::new();
2u64.encode(&mut value, NEW).unwrap();
0x26u64.encode(&mut value, NEW).unwrap(); vec![0xAAu8, 0xBB, 0xCC].encode(&mut value, NEW).unwrap();
1u64.encode(&mut value, NEW).unwrap(); vec![0xDDu8].encode(&mut value, NEW).unwrap();
let mut buf = Vec::new();
value.encode(&mut buf, NEW).unwrap();
let mut bytes = bytes::Bytes::from(buf);
let fill = Fill::param_decode(&mut bytes, NEW).expect("decode");
assert_eq!(fill.filter, None, "no Location Filter in the scope means inherit");
assert!(fill.range_filters, "a Range Filter's presence must be recorded");
}
#[test]
fn skips_allowed_parameters_it_ignores() {
let mut value = Vec::new();
2u64.encode(&mut value, NEW).unwrap();
0x20u64.encode(&mut value, NEW).unwrap(); 42u64.encode(&mut value, NEW).unwrap();
1u64.encode(&mut value, NEW).unwrap(); Filter::Relative(2).param_encode(&mut value, NEW).unwrap();
let mut buf = Vec::new();
value.encode(&mut buf, NEW).unwrap();
let mut bytes = bytes::Bytes::from(buf);
let fill = Fill::param_decode(&mut bytes, NEW).expect("decode");
assert_eq!(fill.filter, Some(Filter::Relative(2)));
}
}