1use std::{fmt, str::FromStr};
2
3use serde::{Deserialize, Serialize};
4
5use crate::{
6 error::{MemoryError, MemoryResult},
7 resource::MemoryResource,
8};
9
10const AUTO_AVAILABLE_FRACTION: f64 = 0.80;
11
12#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
14#[serde(rename_all = "snake_case")]
15pub enum MemoryBudget {
16 #[default]
18 Auto,
19 Bytes(u64),
21 PercentTotal(f64),
23 PercentAvailable(f64),
25}
26
27impl MemoryBudget {
28 pub const fn bytes(bytes: u64) -> Self {
30 Self::Bytes(bytes)
31 }
32
33 pub fn percent_total(percent: f64) -> MemoryResult<Self> {
39 Ok(Self::PercentTotal(validate_percent(percent)? / 100.0))
40 }
41
42 pub fn percent_available(percent: f64) -> MemoryResult<Self> {
48 Ok(Self::PercentAvailable(validate_percent(percent)? / 100.0))
49 }
50
51 pub fn resolve(self, resource: &MemoryResource) -> MemoryResult<u64> {
58 let resolved = match self {
59 Self::Auto => resource
60 .available_bytes
61 .map(|bytes| scaled_bytes(bytes, AUTO_AVAILABLE_FRACTION))
62 .or(resource.total_bytes.map(|bytes| scaled_bytes(bytes, 0.5)))
63 .ok_or_else(|| MemoryError::UnknownCapacity {
64 resource: resource.name.clone(),
65 budget: self,
66 basis: "available",
67 })?,
68 Self::Bytes(bytes) => bytes,
69 Self::PercentTotal(fraction) => {
70 validate_fraction(fraction)?;
71 scaled_bytes(
72 resource
73 .total_bytes
74 .ok_or_else(|| MemoryError::UnknownCapacity {
75 resource: resource.name.clone(),
76 budget: self,
77 basis: "total",
78 })?,
79 fraction,
80 )
81 }
82 Self::PercentAvailable(fraction) => {
83 validate_fraction(fraction)?;
84 scaled_bytes(
85 resource
86 .available_bytes
87 .ok_or_else(|| MemoryError::UnknownCapacity {
88 resource: resource.name.clone(),
89 budget: self,
90 basis: "available",
91 })?,
92 fraction,
93 )
94 }
95 };
96 if resolved == 0 {
97 return Err(MemoryError::InvalidBudget(
98 "resolved budget must be greater than zero".into(),
99 ));
100 }
101 Ok(resolved)
102 }
103}
104
105impl fmt::Display for MemoryBudget {
106 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
107 match self {
108 Self::Auto => formatter.write_str("auto"),
109 Self::Bytes(bytes) => write!(formatter, "{bytes} B"),
110 Self::PercentTotal(value) => write!(formatter, "{}% total", value * 100.0),
111 Self::PercentAvailable(value) => write!(formatter, "{}% available", value * 100.0),
112 }
113 }
114}
115
116impl FromStr for MemoryBudget {
117 type Err = MemoryError;
118
119 fn from_str(input: &str) -> Result<Self, Self::Err> {
120 let normalized = input.trim().to_ascii_lowercase();
121 if normalized == "auto" {
122 return Ok(Self::Auto);
123 }
124 if let Some((percent, suffix)) = normalized.split_once('%') {
125 let percent = percent
126 .trim()
127 .parse::<f64>()
128 .map_err(|_| MemoryError::InvalidBudget(input.into()))?;
129 return match suffix.trim() {
130 "" | "total" => Self::percent_total(percent),
131 "available" | "free" | "remaining" => Self::percent_available(percent),
132 _ => Err(MemoryError::InvalidBudget(input.into())),
133 };
134 }
135 parse_bytes(&normalized).map(Self::Bytes)
136 }
137}
138
139#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
141pub struct MemoryPlan {
142 pub host: MemoryBudget,
144 pub device: Option<MemoryBudget>,
146}
147
148impl Default for MemoryPlan {
149 fn default() -> Self {
150 Self {
151 host: MemoryBudget::Auto,
152 device: Some(MemoryBudget::Auto),
153 }
154 }
155}
156
157impl MemoryPlan {
158 pub const fn host(host: MemoryBudget) -> Self {
160 Self { host, device: None }
161 }
162 pub const fn host_device(host: MemoryBudget, device: MemoryBudget) -> Self {
164 Self {
165 host,
166 device: Some(device),
167 }
168 }
169}
170
171fn validate_percent(percent: f64) -> MemoryResult<f64> {
172 if percent.is_finite() && percent > 0.0 && percent <= 100.0 {
173 Ok(percent)
174 } else {
175 Err(MemoryError::InvalidBudget(
176 "percentage must be finite and in (0, 100]".into(),
177 ))
178 }
179}
180
181fn validate_fraction(fraction: f64) -> MemoryResult<()> {
182 validate_percent(fraction * 100.0).map(|_| ())
183}
184fn scaled_bytes(bytes: u64, fraction: f64) -> u64 {
185 ((bytes as f64) * fraction).floor().min(u64::MAX as f64) as u64
186}
187
188fn parse_bytes(input: &str) -> MemoryResult<u64> {
189 let split = input
190 .find(|character: char| !character.is_ascii_digit() && character != '.')
191 .unwrap_or(input.len());
192 let (number, unit) = input.split_at(split);
193 let multiplier = match unit.trim() {
194 "" | "b" | "byte" | "bytes" => 1,
195 "kb" => 1_000,
196 "mb" => 1_000_000,
197 "gb" => 1_000_000_000,
198 "tb" => 1_000_000_000_000,
199 "kib" => 1 << 10,
200 "mib" => 1 << 20,
201 "gib" => 1 << 30,
202 "tib" => 1 << 40,
203 _ => return Err(MemoryError::InvalidBudget(input.into())),
204 };
205 let (whole, fraction) = match number.split_once('.') {
206 Some((whole, fraction)) if !fraction.contains('.') => (whole, fraction),
207 Some(_) => return Err(MemoryError::InvalidBudget(input.into())),
208 None => (number, ""),
209 };
210 if whole.is_empty() && fraction.is_empty()
211 || !whole.bytes().all(|digit| digit.is_ascii_digit())
212 || !fraction.bytes().all(|digit| digit.is_ascii_digit())
213 || !whole
214 .bytes()
215 .chain(fraction.bytes())
216 .any(|digit| digit != b'0')
217 {
218 return Err(MemoryError::InvalidBudget(input.into()));
219 }
220 let whole = if whole.is_empty() {
221 0
222 } else {
223 whole
224 .parse::<u64>()
225 .map_err(|_| MemoryError::InvalidBudget(input.into()))?
226 };
227 let whole_bytes = whole
228 .checked_mul(multiplier)
229 .ok_or_else(|| MemoryError::InvalidBudget(input.into()))?;
230 let fractional_bytes = fraction.bytes().rev().try_fold(0_u64, |carry, digit| {
231 u64::from(digit - b'0')
232 .checked_mul(multiplier)
233 .and_then(|value| value.checked_add(carry))
234 .map(|value| value / 10)
235 });
236 whole_bytes
237 .checked_add(fractional_bytes.ok_or_else(|| MemoryError::InvalidBudget(input.into()))?)
238 .ok_or_else(|| MemoryError::InvalidBudget(input.into()))
239}
240
241#[cfg(test)]
242mod tests {
243 use super::*;
244 use crate::{CapacitySource, MemoryResourceKind};
245
246 fn resource() -> MemoryResource {
247 MemoryResource {
248 id: "test".into(),
249 name: "Test".into(),
250 kind: MemoryResourceKind::Device,
251 total_bytes: Some(1_000),
252 available_bytes: Some(500),
253 capacity_source: CapacitySource::User,
254 device_identity: None,
255 }
256 }
257
258 #[test]
259 fn parses_and_resolves_budgets() {
260 assert_eq!(
261 "8 GiB".parse(),
262 Ok(MemoryBudget::Bytes(8 * 1024_u64.pow(3)))
263 );
264 assert_eq!("70% total".parse(), Ok(MemoryBudget::PercentTotal(0.7)));
265 assert_eq!(
266 "60% available".parse(),
267 Ok(MemoryBudget::PercentAvailable(0.6))
268 );
269 assert_eq!(MemoryBudget::Auto.resolve(&resource()), Ok(400));
270 assert_eq!(
271 MemoryBudget::PercentTotal(0.5).resolve(&resource()),
272 Ok(500)
273 );
274 assert_eq!(
275 MemoryBudget::PercentAvailable(0.5).resolve(&resource()),
276 Ok(250)
277 );
278 }
279
280 #[test]
281 fn parses_absolute_budgets_exactly_at_boundaries() {
282 for (unit, expected) in [
283 ("b", 1),
284 ("byte", 1),
285 ("bytes", 1),
286 ("kb", 1_000),
287 ("mb", 1_000_000),
288 ("gb", 1_000_000_000),
289 ("tb", 1_000_000_000_000),
290 ("kib", 1 << 10),
291 ("mib", 1 << 20),
292 ("gib", 1 << 30),
293 ("tib", 1 << 40),
294 ] {
295 assert_eq!(
296 format!("1 {unit}").parse(),
297 Ok(MemoryBudget::Bytes(expected))
298 );
299 }
300 assert_eq!("1.5 KiB".parse(), Ok(MemoryBudget::Bytes(1_536)));
301 assert_eq!(".5 kb".parse(), Ok(MemoryBudget::Bytes(500)));
302 assert_eq!("1.999 B".parse(), Ok(MemoryBudget::Bytes(1)));
303 assert_eq!("0.1 B".parse(), Ok(MemoryBudget::Bytes(0)));
304 assert_eq!(
305 "9007199254740991 B".parse(),
306 Ok(MemoryBudget::Bytes(9_007_199_254_740_991))
307 );
308 assert_eq!(
309 "9007199254740993 B".parse(),
310 Ok(MemoryBudget::Bytes(9_007_199_254_740_993))
311 );
312 assert_eq!(
313 "18446744073709551615 B".parse(),
314 Ok(MemoryBudget::Bytes(u64::MAX))
315 );
316 assert_eq!(
317 "18446744073709551.615 KB".parse(),
318 Ok(MemoryBudget::Bytes(u64::MAX))
319 );
320 assert_eq!(
321 MemoryBudget::Bytes(u64::MAX).to_string().parse(),
322 Ok(MemoryBudget::Bytes(u64::MAX))
323 );
324 }
325
326 #[test]
327 fn rejects_invalid_budgets() {
328 for input in [
329 "",
330 "0",
331 ".",
332 "NaN",
333 "inf",
334 "1.2.3 B",
335 "1 XB",
336 "bytes",
337 "18446744073709551616 B",
338 "18446744073709551.616 KB",
339 ] {
340 assert!(
341 input.parse::<MemoryBudget>().is_err(),
342 "{input:?} should be invalid"
343 );
344 }
345 for percent in [0.0, -1.0, 100.1, f64::NAN, f64::INFINITY] {
346 assert!(MemoryBudget::percent_total(percent).is_err());
347 assert!(MemoryBudget::percent_available(percent).is_err());
348 }
349 assert_eq!(
350 MemoryBudget::percent_total(100.0),
351 Ok(MemoryBudget::PercentTotal(1.0))
352 );
353 }
354}