use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(transparent)]
pub struct ProtocolVersion(pub u32);
impl ProtocolVersion {
pub const fn new(version: u32) -> Self {
Self(version)
}
pub const fn get(self) -> u32 {
self.0
}
}
impl From<u32> for ProtocolVersion {
fn from(value: u32) -> Self {
Self(value)
}
}
impl std::fmt::Display for ProtocolVersion {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.0)
}
}
pub const CURRENT_VERSION: ProtocolVersion = ProtocolVersion(1);
const _: () = assert!(CURRENT_VERSION.0 >= 1, "CURRENT_VERSION must be >= 1");
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SupportedVersions {
min: ProtocolVersion,
max: ProtocolVersion,
}
impl SupportedVersions {
pub const fn current() -> Self {
let max = CURRENT_VERSION.0;
let min = if max > 1 { max - 1 } else { 1 };
Self {
min: ProtocolVersion(min),
max: ProtocolVersion(max),
}
}
pub const fn new(
min: ProtocolVersion,
max: ProtocolVersion,
) -> Result<Self, SupportedVersionsError> {
if min.0 < 1 {
return Err(SupportedVersionsError::MinVersionZero);
}
if min.0 > max.0 {
return Err(SupportedVersionsError::InvertedRange);
}
if max.0 > CURRENT_VERSION.0 {
return Err(SupportedVersionsError::MaxAboveCurrent);
}
Ok(Self { min, max })
}
pub const fn min(self) -> ProtocolVersion {
self.min
}
pub const fn max(self) -> ProtocolVersion {
self.max
}
pub const fn contains(self, version: ProtocolVersion) -> bool {
version.0 >= self.min.0 && version.0 <= self.max.0
}
}
impl Default for SupportedVersions {
fn default() -> Self {
Self::current()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub enum SupportedVersionsError {
#[error("supported-version range floor must be >= 1 (version 0 does not exist)")]
MinVersionZero,
#[error("supported-version range is inverted: min must not exceed max")]
InvertedRange,
#[error("supported-version range max must not exceed the current protocol version")]
MaxAboveCurrent,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_accepts_the_current_range() {
let supported =
SupportedVersions::new(ProtocolVersion::new(1), ProtocolVersion::new(1)).unwrap();
assert_eq!(supported.min(), ProtocolVersion::new(1));
assert_eq!(supported.max(), ProtocolVersion::new(1));
}
#[test]
fn new_rejects_a_range_above_current_version() {
for (min, max) in [(1, 2), (2, 2)] {
let err = SupportedVersions::new(ProtocolVersion::new(min), ProtocolVersion::new(max))
.unwrap_err();
assert_eq!(err, SupportedVersionsError::MaxAboveCurrent);
}
}
#[test]
fn new_rejects_an_inverted_range() {
let err =
SupportedVersions::new(ProtocolVersion::new(5), ProtocolVersion::new(2)).unwrap_err();
assert_eq!(err, SupportedVersionsError::InvertedRange);
}
#[test]
fn new_rejects_a_zero_min_version() {
let err =
SupportedVersions::new(ProtocolVersion::new(0), ProtocolVersion::new(1)).unwrap_err();
assert_eq!(err, SupportedVersionsError::MinVersionZero);
let err =
SupportedVersions::new(ProtocolVersion::new(0), ProtocolVersion::new(0)).unwrap_err();
assert_eq!(err, SupportedVersionsError::MinVersionZero);
}
#[test]
fn current_version_is_at_least_one() {
assert!(CURRENT_VERSION.get() >= 1);
}
#[test]
fn contains_is_inclusive_at_both_bounds() {
let supported =
SupportedVersions::new(ProtocolVersion::new(1), ProtocolVersion::new(1)).unwrap();
assert!(!supported.contains(ProtocolVersion::new(0)));
assert!(supported.contains(ProtocolVersion::new(1)));
assert!(!supported.contains(ProtocolVersion::new(2)));
}
#[test]
fn current_saturates_at_version_one() {
let supported = SupportedVersions::current();
assert_eq!(supported.min(), ProtocolVersion::new(1));
assert_eq!(supported.max(), ProtocolVersion::new(1));
}
}