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}