use bytes::{Buf, BufMut};
use num_enum::{IntoPrimitive, TryFromPrimitive};
use crate::{Hop, Hops, Path, coding::*, origin::Cost};
use super::{Message, Version, message::decode_size};
const ANNOUNCE_START: u64 = 0;
const ANNOUNCE_END: u64 = 1;
const ANNOUNCE_RESTART: u64 = 2;
pub fn restart_supported(version: Version) -> bool {
!matches!(
version,
Version::Lite01 | Version::Lite02 | Version::Lite03 | Version::Lite04
)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum AnnounceBroadcast<'a> {
Active {
suffix: PathRef<'a>,
hops: HopsRef,
cost: Cost,
},
Ended { suffix: Path<'a>, hops: Hops },
EndedId { id: u64 },
Restart { id: u64, hops: HopsRef, cost: Cost },
Skipped,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PathRef<'a> {
pub base: u64,
pub keep: u64,
pub rest: Path<'a>,
}
impl<'a> PathRef<'a> {
pub fn literal(rest: Path<'a>) -> Self {
Self { base: 0, keep: 0, rest }
}
}
impl Encode<Version> for PathRef<'_> {
fn encode<W: BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if version.has_announce_compression() {
self.base.encode(w, version)?;
self.keep.encode(w, version)?;
} else if self.base != 0 || self.keep != 0 {
return Err(EncodeError::Version);
}
self.rest.encode(w, version)
}
}
impl Decode<Version> for PathRef<'_> {
fn decode<B: Buf>(buf: &mut B, version: Version) -> Result<Self, DecodeError> {
if !version.has_announce_compression() {
return Ok(Self::literal(Path::decode(buf, version)?));
}
let base = u64::decode(buf, version)?;
let keep = u64::decode(buf, version)?;
if base == 0 && keep != 0 {
return Err(DecodeError::InvalidValue);
}
let rest = Path::decode(buf, version)?;
Ok(Self { base, keep, rest })
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct HopsRef {
pub base: u64,
pub literal: Hops,
pub keep: u64,
}
impl HopsRef {
pub fn literal(literal: Hops) -> Self {
Self {
base: 0,
literal,
keep: 0,
}
}
}
impl Encode<Version> for HopsRef {
fn encode<W: BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_announce_compression() {
if self.base != 0 || self.keep != 0 {
return Err(EncodeError::Version);
}
return self.literal.encode(w, version);
}
self.base.encode(w, version)?;
self.literal.encode(w, version)?;
self.keep.encode(w, version)
}
}
impl Decode<Version> for HopsRef {
fn decode<B: Buf>(buf: &mut B, version: Version) -> Result<Self, DecodeError> {
if !version.has_announce_compression() {
return Ok(Self::literal(Hops::decode(buf, version)?));
}
let base = u64::decode(buf, version)?;
let literal = Hops::decode(buf, version)?;
let keep = u64::decode(buf, version)?;
if base == 0 && keep != 0 {
return Err(DecodeError::InvalidValue);
}
Ok(Self { base, literal, keep })
}
}
impl AnnounceBroadcast<'_> {
#[cfg(test)]
pub fn into_owned(self) -> AnnounceBroadcast<'static> {
match self {
Self::Active { suffix, hops, cost } => AnnounceBroadcast::Active {
suffix: PathRef {
base: suffix.base,
keep: suffix.keep,
rest: suffix.rest.into_owned(),
},
hops,
cost,
},
Self::Ended { suffix, hops } => AnnounceBroadcast::Ended {
suffix: suffix.into_owned(),
hops,
},
Self::EndedId { id } => AnnounceBroadcast::EndedId { id },
Self::Restart { id, hops, cost } => AnnounceBroadcast::Restart { id, hops, cost },
Self::Skipped => AnnounceBroadcast::Skipped,
}
}
}
impl Encode<Version> for Cost {
fn encode<W: BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_route_cost() {
return Ok(());
}
self.warm.encode(w, version)?;
self.cold.encode(w, version)
}
}
impl Decode<Version> for Cost {
fn decode<B: Buf>(buf: &mut B, version: Version) -> Result<Self, DecodeError> {
if !version.has_route_cost() {
return Ok(Cost::UNKNOWN);
}
Ok(Cost {
warm: u64::decode(buf, version)?,
cold: u64::decode(buf, version)?,
})
}
}
impl Encode<Version> for AnnounceBroadcast<'_> {
fn encode<W: BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if version.has_announce_id() {
let mut body = Vec::new();
let typ = match self {
Self::Active { suffix, hops, cost } => {
suffix.encode(&mut body, version)?;
hops.encode(&mut body, version)?;
cost.encode(&mut body, version)?;
ANNOUNCE_START
}
Self::EndedId { id } => {
id.encode(&mut body, version)?;
ANNOUNCE_END
}
Self::Restart { id, hops, cost } => {
id.encode(&mut body, version)?;
hops.encode(&mut body, version)?;
cost.encode(&mut body, version)?;
ANNOUNCE_RESTART
}
Self::Ended { .. } => return Err(EncodeError::Version),
Self::Skipped => return Err(EncodeError::Unsupported),
};
typ.encode(w, version)?;
(body.len() as u64).encode(w, version)?;
w.put_slice(&body);
return Ok(());
}
let mut body = Vec::new();
match self {
Self::Active { suffix, hops, .. } => {
if suffix.base != 0 || hops.base != 0 {
return Err(EncodeError::Version);
}
AnnounceStatus::Active.encode(&mut body, version)?;
suffix.rest.encode(&mut body, version)?;
encode_hops(&mut body, version, &hops.literal)?;
}
Self::Ended { suffix, hops } => {
AnnounceStatus::Ended.encode(&mut body, version)?;
suffix.encode(&mut body, version)?;
encode_hops(&mut body, version, hops)?;
}
Self::EndedId { .. } | Self::Restart { .. } | Self::Skipped => {
return Err(EncodeError::Version);
}
}
(body.len() as u64).encode(w, version)?;
w.put_slice(&body);
Ok(())
}
}
impl Decode<Version> for AnnounceBroadcast<'_> {
fn decode<B: Buf>(buf: &mut B, version: Version) -> Result<Self, DecodeError> {
if version.has_announce_id() {
let typ = u64::decode(buf, version)?;
let size = decode_size(buf, version)?;
if buf.remaining() < size {
return Err(DecodeError::Short);
}
let mut body = buf.take(size);
let msg = match typ {
ANNOUNCE_START => Self::Active {
suffix: PathRef::decode(&mut body, version)?,
hops: HopsRef::decode(&mut body, version)?,
cost: Cost::decode(&mut body, version)?,
},
ANNOUNCE_END => Self::EndedId {
id: u64::decode(&mut body, version)?,
},
ANNOUNCE_RESTART => Self::Restart {
id: u64::decode(&mut body, version)?,
hops: HopsRef::decode(&mut body, version)?,
cost: Cost::decode(&mut body, version)?,
},
_ => {
let remaining = body.remaining();
bytes::Buf::advance(&mut body, remaining);
Self::Skipped
}
};
if body.remaining() > 0 {
return Err(DecodeError::Long);
}
return Ok(msg);
}
let size = decode_size(buf, version)?;
if buf.remaining() < size {
return Err(DecodeError::Short);
}
let mut body = buf.take(size);
let msg = Self::decode_legacy(&mut body, version)?;
if body.remaining() > 0 {
return Err(DecodeError::Long);
}
Ok(msg)
}
}
impl AnnounceBroadcast<'_> {
fn decode_legacy<R: Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let status = AnnounceStatus::decode(r, version)?;
let suffix = Path::decode(r, version)?;
let hops = match version {
Version::Lite01 | Version::Lite02 => Hops::new(),
Version::Lite03 => {
let count = u64::decode(r, version)? as usize;
let mut list = Hops::new();
for _ in 0..count {
list.push(Hop::UNKNOWN)?;
}
list
}
_ => Hops::decode(r, version)?,
};
Ok(match status {
AnnounceStatus::Active => Self::Active {
suffix: PathRef::literal(suffix),
hops: HopsRef::literal(hops),
cost: Cost::UNKNOWN,
},
AnnounceStatus::Ended => Self::Ended { suffix, hops },
AnnounceStatus::Restart if restart_supported(version) => Self::Active {
suffix: PathRef::literal(suffix),
hops: HopsRef::literal(hops),
cost: Cost::UNKNOWN,
},
AnnounceStatus::Restart => return Err(DecodeError::InvalidValue),
})
}
}
fn encode_hops<W: bytes::BufMut>(w: &mut W, version: Version, hops: &Hops) -> Result<(), EncodeError> {
match version {
Version::Lite01 | Version::Lite02 => Ok(()),
Version::Lite03 => (hops.len() as u64).encode(w, version),
_ => hops.encode(w, version),
}
}
#[derive(Clone, Debug)]
pub struct AnnounceRequest<'a> {
pub prefix: Path<'a>,
pub exclude_hop: u64,
pub hidden: bool,
}
impl Message for AnnounceRequest<'_> {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let prefix = Path::decode(r, version)?;
let exclude_hop = match version.has_exclude_hop() {
true => u64::decode(r, version)?,
false => 0,
};
let hidden = match version.has_hidden() {
true => bool::decode(r, version)?,
false => false,
};
Ok(Self {
prefix,
exclude_hop,
hidden,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
self.prefix.encode(w, version)?;
if version.has_exclude_hop() {
self.exclude_hop.encode(w, version)?;
}
if version.has_hidden() {
self.hidden.encode(w, version)?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, IntoPrimitive, TryFromPrimitive)]
#[repr(u8)]
enum AnnounceStatus {
Ended = 0,
Active = 1,
Restart = 2,
}
impl Decode<Version> for AnnounceStatus {
fn decode<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
let status = u8::decode(r, version)?;
status.try_into().map_err(|_| DecodeError::InvalidValue)
}
}
impl Encode<Version> for AnnounceStatus {
fn encode<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
(*self as u8).encode(w, version)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AnnounceInit<'a> {
pub suffixes: Vec<Path<'a>>,
}
impl Message for AnnounceInit<'_> {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
match version {
Version::Lite01 | Version::Lite02 => {}
_ => {
return Err(DecodeError::Version);
}
}
let count = u64::decode(r, version)?;
let mut paths = Vec::with_capacity(count.min(1024) as usize);
for _ in 0..count {
paths.push(Path::decode(r, version)?);
}
Ok(Self { suffixes: paths })
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
match version {
Version::Lite01 | Version::Lite02 => {}
_ => {
return Err(EncodeError::Version);
}
}
(self.suffixes.len() as u64).encode(w, version)?;
for path in &self.suffixes {
path.encode(w, version)?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct AnnounceOk {
pub origin: Hop,
pub active: u64,
}
impl Message for AnnounceOk {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
if !version.has_announce_ok() {
return Err(DecodeError::Version);
}
let origin = Hop::decode(r, version)?;
let active = u64::decode(r, version)?;
Ok(Self { origin, active })
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_announce_ok() {
return Err(EncodeError::Version);
}
self.origin.encode(w, version)?;
self.active.encode(w, version)
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Buf;
fn encode_forged_restart(version: Version) -> bytes::Bytes {
let mut buf = bytes::BytesMut::new();
AnnounceBroadcast::Active {
suffix: PathRef::literal(Path::new("foo/bar")),
hops: HopsRef::default(),
cost: Cost::default(),
}
.encode(&mut buf, version)
.expect("encode");
assert_eq!(
buf[1],
u8::from(AnnounceStatus::Active),
"expected an Active status byte"
);
buf[1] = u8::from(AnnounceStatus::Restart);
buf.freeze()
}
#[test]
fn decodes_explicit_restart_status_as_active_on_lite05() {
let version = Version::Lite05;
let mut slice = encode_forged_restart(version);
let decoded = AnnounceBroadcast::decode(&mut slice, version).expect("explicit restart must decode");
assert!(!slice.has_remaining(), "trailing bytes after decode");
assert!(
matches!(decoded, AnnounceBroadcast::Active { .. }),
"restart should decode as Active"
);
}
#[test]
fn rejects_explicit_restart_status_before_lite05() {
let version = Version::Lite04;
let mut slice = encode_forged_restart(version);
assert!(
matches!(
AnnounceBroadcast::decode(&mut slice, version),
Err(DecodeError::InvalidValue)
),
"restart status must be rejected before lite-05"
);
}
fn round_trip(msg: &AnnounceOk) -> AnnounceOk {
let mut buf = bytes::BytesMut::new();
msg.encode(&mut buf, Version::Lite05).unwrap();
let mut slice = &buf[..];
let got = AnnounceOk::decode(&mut slice, Version::Lite05).unwrap();
assert!(slice.is_empty(), "trailing bytes after decode");
got
}
#[test]
fn announce_ok_round_trip() {
let msg = AnnounceOk {
origin: Hop::new(42).unwrap(),
active: 3,
};
assert_eq!(round_trip(&msg), msg);
}
#[test]
fn announce_ok_zero_active() {
let msg = AnnounceOk {
origin: Hop::new(7).unwrap(),
active: 0,
};
assert_eq!(round_trip(&msg), msg);
}
fn broadcast_round_trip(msg: &AnnounceBroadcast, version: Version) -> AnnounceBroadcast<'static> {
let mut buf = bytes::BytesMut::new();
msg.encode(&mut buf, version).unwrap();
let mut slice = &buf[..];
let got = AnnounceBroadcast::decode(&mut slice, version).unwrap();
assert!(slice.is_empty(), "trailing bytes after decode");
got.into_owned()
}
#[test]
fn announce_broadcast_round_trip_on_lite05() {
let mut hops = Hops::new();
hops.push(Hop::new(7).unwrap()).unwrap();
let msg = AnnounceBroadcast::Active {
suffix: PathRef::literal(Path::new("room/cam")),
hops: HopsRef::literal(hops.clone()),
cost: Cost::UNKNOWN,
};
assert_eq!(broadcast_round_trip(&msg, Version::Lite05), msg);
let ended = AnnounceBroadcast::Ended {
suffix: Path::new("room/cam"),
hops: Hops::new(),
};
assert_eq!(broadcast_round_trip(&ended, Version::Lite05), ended);
}
#[test]
fn announce_broadcast_round_trip_on_lite06() {
let mut hops = Hops::new();
hops.push(Hop::new(7).unwrap()).unwrap();
let cost = Cost { warm: 12, cold: 30 };
let active = AnnounceBroadcast::Active {
suffix: PathRef::literal(Path::new("room/cam")),
hops: HopsRef::literal(hops.clone()),
cost,
};
assert_eq!(broadcast_round_trip(&active, Version::Lite06), active);
let ended = AnnounceBroadcast::EndedId { id: 3 };
assert_eq!(broadcast_round_trip(&ended, Version::Lite06), ended);
let restart = AnnounceBroadcast::Restart {
id: 3,
hops: HopsRef::literal(hops),
cost,
};
assert_eq!(broadcast_round_trip(&restart, Version::Lite06), restart);
}
#[test]
fn announce_broadcast_round_trip_on_lite07() {
let mut hops = Hops::new();
hops.push(Hop::new(7).unwrap()).unwrap();
let cost = Cost { warm: 12, cold: 30 };
let active = AnnounceBroadcast::Active {
suffix: PathRef {
base: 2,
keep: 3,
rest: Path::new("cam"),
},
hops: HopsRef {
base: 1,
literal: hops.clone(),
keep: 2,
},
cost,
};
assert_eq!(broadcast_round_trip(&active, Version::Lite07), active);
let restart = AnnounceBroadcast::Restart {
id: 3,
hops: HopsRef {
base: 4,
literal: hops,
keep: 1,
},
cost,
};
assert_eq!(broadcast_round_trip(&restart, Version::Lite07), restart);
}
#[test]
fn a_keep_without_a_base_is_rejected() {
for (path_keep, hop_keep) in [(1u8, 0u8), (0, 1)] {
let body = [0, path_keep, 0, 0, 0, hop_keep, 0, 0];
let mut buf = vec![ANNOUNCE_START as u8, body.len() as u8];
buf.extend_from_slice(&body);
assert!(matches!(
AnnounceBroadcast::decode(&mut &buf[..], Version::Lite07),
Err(DecodeError::InvalidValue)
));
}
}
#[test]
fn a_base_needs_lite07() {
let msg = AnnounceBroadcast::Active {
suffix: PathRef {
base: 1,
keep: 1,
rest: Path::new("cam"),
},
hops: HopsRef::default(),
cost: Cost::default(),
};
for version in [Version::Lite05, Version::Lite06] {
let mut buf = bytes::BytesMut::new();
assert!(matches!(msg.encode(&mut buf, version), Err(EncodeError::Version)));
}
}
#[test]
fn announce_broadcast_rejects_cross_version_forms() {
let mut buf = bytes::BytesMut::new();
assert!(matches!(
AnnounceBroadcast::EndedId { id: 1 }.encode(&mut buf, Version::Lite05),
Err(EncodeError::Version)
));
assert!(matches!(
AnnounceBroadcast::Restart {
id: 1,
hops: HopsRef::default(),
cost: Cost::default()
}
.encode(&mut buf, Version::Lite05),
Err(EncodeError::Version)
));
assert!(matches!(
AnnounceBroadcast::Ended {
suffix: Path::new("room/cam"),
hops: Hops::new()
}
.encode(&mut buf, Version::Lite06),
Err(EncodeError::Version)
));
}
#[test]
fn route_cost_is_dropped_before_lite06() {
let msg = AnnounceBroadcast::Active {
suffix: PathRef::literal(Path::new("room/cam")),
hops: HopsRef::default(),
cost: Cost { warm: 9, cold: 9 },
};
let got = broadcast_round_trip(&msg, Version::Lite05);
assert_eq!(
got,
AnnounceBroadcast::Active {
suffix: PathRef::literal(Path::new("room/cam")),
hops: HopsRef::default(),
cost: Cost::UNKNOWN,
}
);
}
#[test]
fn charged_cost_stays_encodable() {
let mut buf = Vec::new();
crate::origin::Cost::MAX
.charged(1)
.encode(&mut buf, Version::Lite06)
.expect("a charged cost must stay encodable");
}
#[test]
fn unknown_announce_type_is_skipped() {
let mut body = Vec::new();
Path::new("room/cam").encode(&mut body, Version::Lite06).unwrap();
Hops::new().encode(&mut body, Version::Lite06).unwrap();
Cost::default().encode(&mut body, Version::Lite06).unwrap();
let mut buf = bytes::BytesMut::new();
4u64.encode(&mut buf, Version::Lite06).unwrap();
(body.len() as u64).encode(&mut buf, Version::Lite06).unwrap();
buf.extend_from_slice(&body);
let mut slice = &buf[..];
let got =
AnnounceBroadcast::decode(&mut slice, Version::Lite06).expect("unknown type must not kill the stream");
assert!(slice.is_empty());
assert_eq!(got, AnnounceBroadcast::Skipped);
}
#[test]
fn ended_by_id_is_three_bytes() {
let mut buf = bytes::BytesMut::new();
AnnounceBroadcast::EndedId { id: 42 }
.encode(&mut buf, Version::Lite06)
.unwrap();
assert_eq!(buf.len(), 3);
}
fn request_round_trip(msg: &AnnounceRequest, version: Version) -> AnnounceRequest<'static> {
let mut buf = bytes::BytesMut::new();
msg.encode(&mut buf, version).unwrap();
let mut slice = &buf[..];
let got = AnnounceRequest::decode(&mut slice, version).unwrap();
assert!(slice.is_empty(), "trailing bytes after decode");
AnnounceRequest {
prefix: got.prefix.to_owned(),
exclude_hop: got.exclude_hop,
hidden: got.hidden,
}
}
#[test]
fn announce_request_carries_hidden_from_lite07() {
for hidden in [false, true] {
let msg = AnnounceRequest {
prefix: Path::new("room/"),
exclude_hop: 0,
hidden,
};
assert_eq!(request_round_trip(&msg, Version::Lite07).hidden, hidden);
assert!(!request_round_trip(&msg, Version::Lite06).hidden);
}
}
#[test]
fn announce_request_rejects_a_bad_hidden_flag() {
let mut buf = bytes::BytesMut::new();
let mut body = Vec::new();
Path::new("room").encode(&mut body, Version::Lite07).unwrap();
body.push(2);
(body.len() as u64).encode(&mut buf, Version::Lite07).unwrap();
buf.extend_from_slice(&body);
assert!(AnnounceRequest::decode(&mut &buf[..], Version::Lite07).is_err());
}
#[test]
fn announce_request_carries_exclude_hop_on_lite05() {
let msg = AnnounceRequest {
prefix: Path::new("room/"),
exclude_hop: 42,
hidden: false,
};
assert_eq!(request_round_trip(&msg, Version::Lite05).exclude_hop, 42);
}
#[test]
fn announce_request_drops_exclude_hop_on_lite06() {
let msg = AnnounceRequest {
prefix: Path::new("room/"),
exclude_hop: 42,
hidden: false,
};
assert_eq!(request_round_trip(&msg, Version::Lite06).exclude_hop, 0);
let mut with = bytes::BytesMut::new();
msg.encode(&mut with, Version::Lite05).unwrap();
let mut without = bytes::BytesMut::new();
msg.encode(&mut without, Version::Lite06).unwrap();
assert!(
without.len() < with.len(),
"lite06 must not encode the exclude_hop varint"
);
}
#[test]
fn announce_ok_rejects_old_versions() {
let msg = AnnounceOk {
origin: Hop::new(1).unwrap(),
active: 0,
};
let mut buf = bytes::BytesMut::new();
assert!(matches!(
msg.encode(&mut buf, Version::Lite04),
Err(EncodeError::Version)
));
}
#[test]
fn announce_ok_accepts_zero_origin() {
let mut buf = bytes::BytesMut::new();
AnnounceOk {
origin: Hop::new(1).unwrap(),
active: 0,
}
.encode(&mut buf, Version::Lite05)
.unwrap();
let bytes = &buf[..];
let mut patched = bytes.to_vec();
patched[1] = 0x00;
let mut slice = &patched[..];
let got = AnnounceOk::decode(&mut slice, Version::Lite05).unwrap();
assert_eq!(got.origin.id(), 0);
assert_eq!(got.active, 0);
}
}