use crate::coding::*;
use super::{Message, Parameters, Version};
const PARAM_PROBE: u64 = 0x1;
const PARAM_PATH: u64 = 0x2;
const PARAM_ROLE: u64 = 0x3;
const PARAM_COST: u64 = 0x4;
pub const DEFAULT_COST: u64 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
pub enum ProbeLevel {
#[default]
None,
Report,
Increase,
}
impl ProbeLevel {
fn from_code(code: u64) -> Self {
match code {
0 => Self::None,
1 => Self::Report,
_ => Self::Increase,
}
}
fn to_code(self) -> u64 {
match self {
Self::None => 0,
Self::Report => 1,
Self::Increase => 2,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Role {
Publisher,
Subscriber,
}
impl Role {
fn from_code(code: u64) -> Option<Self> {
match code {
1 => Some(Role::Publisher),
2 => Some(Role::Subscriber),
_ => None,
}
}
fn to_code(self) -> u64 {
match self {
Role::Publisher => 1,
Role::Subscriber => 2,
}
}
pub(crate) fn from_origins(publishes: bool, consumes: bool) -> Option<Self> {
match (publishes, consumes) {
(true, false) => Some(Role::Publisher),
(false, true) => Some(Role::Subscriber),
_ => None,
}
}
pub fn as_str(self) -> &'static str {
match self {
Role::Publisher => "publisher",
Role::Subscriber => "subscriber",
}
}
}
impl std::fmt::Display for Role {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Setup {
pub probe: ProbeLevel,
pub path: Option<String>,
pub role: Option<Role>,
pub cost: Option<u64>,
}
impl Message for Setup {
fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
if !version.has_setup_stream() {
return Err(DecodeError::Version);
}
let params = Parameters::decode(r, version)?;
let probe = params
.get_varint(PARAM_PROBE)?
.map(ProbeLevel::from_code)
.unwrap_or_default();
let path = match params.get_bytes(PARAM_PATH) {
Some(bytes) => Some(
std::str::from_utf8(bytes)
.map_err(|_| DecodeError::InvalidValue)?
.to_string(),
),
None => None,
};
let role = params.get_varint(PARAM_ROLE)?.and_then(Role::from_code);
let cost = params.get_varint(PARAM_COST)?;
Ok(Self {
probe,
path,
role,
cost,
})
}
fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
if !version.has_setup_stream() {
return Err(EncodeError::Version);
}
let mut params = Parameters::default();
if self.probe != ProbeLevel::None {
params.set_varint(PARAM_PROBE, self.probe.to_code());
}
if let Some(path) = &self.path {
params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
}
if let Some(role) = self.role {
params.set_varint(PARAM_ROLE, role.to_code());
}
if let Some(cost) = self.cost {
params.set_varint(PARAM_COST, cost);
}
params.encode(w, version)
}
}
#[derive(Clone, Default)]
pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
impl PeerSetup {
pub fn set(&self, setup: Setup) {
*self.0.lock() = Some(setup);
}
pub async fn probe_level(&self) -> ProbeLevel {
self.wait(|setup| setup.probe).await
}
pub async fn cost(&self) -> Option<u64> {
self.wait(|setup| setup.cost).await
}
async fn wait<T>(&self, f: impl FnOnce(&Setup) -> T) -> T {
let slot = self
.0
.wait(|setup| {
if setup.is_some() {
std::task::Poll::Ready(())
} else {
std::task::Poll::Pending
}
})
.await;
f(slot.as_ref().expect("waited for Some"))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn round_trip(msg: &Setup) -> Setup {
let mut buf = bytes::BytesMut::new();
msg.encode(&mut buf, Version::Lite05).unwrap();
let mut slice = &buf[..];
let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
got
}
#[test]
fn empty_round_trip() {
let msg = Setup::default();
assert_eq!(round_trip(&msg), msg);
}
#[test]
fn probe_levels_round_trip() {
for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
let msg = Setup {
probe,
..Default::default()
};
assert_eq!(round_trip(&msg), msg);
}
}
#[test]
fn cost_round_trip() {
for cost in [None, Some(0), Some(1), Some(7)] {
let msg = Setup {
cost,
..Default::default()
};
assert_eq!(round_trip(&msg), msg);
}
}
#[test]
fn path_round_trip() {
let msg = Setup {
probe: ProbeLevel::Report,
path: Some("/room/123".to_string()),
role: None,
cost: None,
};
assert_eq!(round_trip(&msg), msg);
}
#[test]
fn empty_path_round_trips() {
let msg = Setup {
path: Some(String::new()),
..Default::default()
};
assert_eq!(round_trip(&msg), msg);
}
#[test]
fn roles_round_trip() {
for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
let msg = Setup {
path: Some("/room/123".to_string()),
role,
..Default::default()
};
assert_eq!(round_trip(&msg), msg);
}
}
#[test]
fn unknown_probe_level_saturates_to_increase() {
let mut params = Parameters::default();
params.set_varint(PARAM_PROBE, 99);
let mut body = Vec::new();
params.encode(&mut body, Version::Lite05).unwrap();
let mut buf = bytes::BytesMut::new();
body.len().encode(&mut buf, Version::Lite05).unwrap();
buf.extend_from_slice(&body);
let mut slice = &buf[..];
let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
assert_eq!(got.probe, ProbeLevel::Increase);
}
#[test]
fn role_wire_codes() {
for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
assert_eq!(role.to_code(), code);
assert_eq!(Role::from_code(code), Some(role));
}
}
#[test]
fn unknown_role_decodes_as_bidirectional() {
for code in [0u64, 9, 250] {
let mut params = Parameters::default();
params.set_varint(PARAM_ROLE, code);
let mut body = Vec::new();
params.encode(&mut body, Version::Lite05).unwrap();
let mut buf = bytes::BytesMut::new();
body.len().encode(&mut buf, Version::Lite05).unwrap();
buf.extend_from_slice(&body);
let mut slice = &buf[..];
let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
}
}
#[test]
fn rejects_before_lite05() {
let msg = Setup::default();
let mut buf = bytes::BytesMut::new();
assert!(matches!(
msg.encode(&mut buf, Version::Lite04),
Err(EncodeError::Version)
));
}
#[test]
fn ignores_unknown_parameters() {
let mut params = Parameters::default();
params.set_bytes(PARAM_PATH, b"/foo".to_vec());
params.set_bytes(0xbeef, b"whatever".to_vec());
let mut body = Vec::new();
params.encode(&mut body, Version::Lite05).unwrap();
let mut buf = bytes::BytesMut::new();
body.len().encode(&mut buf, Version::Lite05).unwrap();
buf.extend_from_slice(&body);
let mut slice = &buf[..];
let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
assert_eq!(got.path.as_deref(), Some("/foo"));
}
}