Skip to main content

pbox_core/
vmid.rs

1use serde::de::Error as DeError;
2use serde::{Deserialize, Deserializer, Serialize, Serializer};
3use std::borrow::Borrow;
4use std::fmt;
5use std::str::FromStr;
6use thiserror::Error;
7
8#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct VmidPattern {
10    prefix: u64,
11    wildcard_digits: u32,
12    first: u64,
13    last: u64,
14}
15
16#[derive(Debug, Error, Clone, PartialEq, Eq)]
17pub enum VmidError {
18    #[error("VMID pattern must match ^[1-9][0-9]*x+$")]
19    InvalidPattern,
20    #[error("VMID pattern is too large for a 64-bit VMID")]
21    Overflow,
22    #[error("all VMIDs in pattern {0} are already in use")]
23    Exhausted(String),
24}
25
26impl VmidPattern {
27    pub fn parse(pattern: &str) -> Result<Self, VmidError> {
28        let bytes = pattern.as_bytes();
29        if bytes.len() < 2 || !bytes[0].is_ascii_digit() || bytes[0] == b'0' {
30            return Err(VmidError::InvalidPattern);
31        }
32
33        let wildcard_start = bytes
34            .iter()
35            .position(|byte| *byte == b'x')
36            .ok_or(VmidError::InvalidPattern)?;
37        if wildcard_start == 0
38            || bytes[wildcard_start..].iter().any(|byte| *byte != b'x')
39            || bytes[..wildcard_start]
40                .iter()
41                .any(|byte| !byte.is_ascii_digit())
42        {
43            return Err(VmidError::InvalidPattern);
44        }
45
46        let prefix_text = &pattern[..wildcard_start];
47        if prefix_text.len() > 1 && prefix_text.starts_with('0') {
48            return Err(VmidError::InvalidPattern);
49        }
50        let prefix = prefix_text
51            .parse::<u64>()
52            .map_err(|_| VmidError::Overflow)?;
53        let wildcard_digits =
54            u32::try_from(bytes.len() - wildcard_start).map_err(|_| VmidError::Overflow)?;
55        let factor = 10u64
56            .checked_pow(wildcard_digits)
57            .ok_or(VmidError::Overflow)?;
58        let first = prefix.checked_mul(factor).ok_or(VmidError::Overflow)?;
59        let last = first.checked_add(factor - 1).ok_or(VmidError::Overflow)?;
60
61        Ok(Self {
62            prefix,
63            wildcard_digits,
64            first,
65            last,
66        })
67    }
68
69    pub fn prefix(&self) -> u64 {
70        self.prefix
71    }
72
73    pub fn wildcard_digits(&self) -> u32 {
74        self.wildcard_digits
75    }
76
77    pub fn first(&self) -> u64 {
78        self.first
79    }
80
81    pub fn last(&self) -> u64 {
82        self.last
83    }
84
85    pub fn contains(&self, vmid: u64) -> bool {
86        (self.first..=self.last).contains(&vmid)
87    }
88
89    pub fn allocate_lowest<I>(&self, used: I) -> Result<u64, VmidError>
90    where
91        I: IntoIterator,
92        I::Item: Borrow<u64>,
93    {
94        let used: std::collections::BTreeSet<u64> =
95            used.into_iter().map(|value| *value.borrow()).collect();
96        let mut candidate = self.first;
97        loop {
98            if !used.contains(&candidate) {
99                return Ok(candidate);
100            }
101            if candidate == self.last {
102                return Err(VmidError::Exhausted(self.to_string()));
103            }
104            candidate += 1;
105        }
106    }
107}
108
109impl fmt::Display for VmidPattern {
110    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
111        write!(
112            formatter,
113            "{}{}",
114            self.prefix,
115            "x".repeat(self.wildcard_digits as usize)
116        )
117    }
118}
119
120impl FromStr for VmidPattern {
121    type Err = VmidError;
122
123    fn from_str(value: &str) -> Result<Self, Self::Err> {
124        Self::parse(value)
125    }
126}
127
128impl Serialize for VmidPattern {
129    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
130    where
131        S: Serializer,
132    {
133        serializer.serialize_str(&self.to_string())
134    }
135}
136
137impl<'de> Deserialize<'de> for VmidPattern {
138    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
139    where
140        D: Deserializer<'de>,
141    {
142        let value = String::deserialize(deserializer)?;
143        Self::parse(&value).map_err(D::Error::custom)
144    }
145}
146
147#[cfg(test)]
148mod tests {
149    use super::*;
150
151    #[test]
152    fn pattern_exposes_expected_range() {
153        let pattern = VmidPattern::parse("9xxx").expect("valid pattern");
154        assert_eq!(pattern.first(), 9000);
155        assert_eq!(pattern.last(), 9999);
156        assert!(pattern.contains(9000));
157        assert!(!pattern.contains(8999));
158    }
159
160    #[test]
161    fn allocation_returns_lowest_unused_candidate() {
162        let pattern = VmidPattern::parse("42xx").expect("valid pattern");
163        assert_eq!(pattern.allocate_lowest([4202, 4200, 4201]), Ok(4203));
164    }
165
166    #[test]
167    fn invalid_patterns_and_exhaustion_are_reported() {
168        for value in ["0xxx", "9", "9xx9", "xx9", "09xx"] {
169            assert!(VmidPattern::parse(value).is_err(), "accepted {value}");
170        }
171        let pattern = VmidPattern::parse("1x").expect("valid pattern");
172        assert!(matches!(
173            pattern.allocate_lowest([10, 11, 12, 13, 14, 15, 16, 17, 18, 19]),
174            Err(VmidError::Exhausted(_))
175        ));
176    }
177}