Skip to main content

laddu_memory/
budget.rs

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/// A requested memory limit.
13#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
14#[serde(rename_all = "snake_case")]
15pub enum MemoryBudget {
16    /// Automatically use 80% of currently available memory.
17    #[default]
18    Auto,
19    /// An absolute number of bytes.
20    Bytes(u64),
21    /// A fraction in `(0, 1]` of total physical capacity.
22    PercentTotal(f64),
23    /// A fraction in `(0, 1]` of currently available capacity.
24    PercentAvailable(f64),
25}
26
27impl MemoryBudget {
28    /// Creates an absolute byte budget.
29    pub const fn bytes(bytes: u64) -> Self {
30        Self::Bytes(bytes)
31    }
32
33    /// Creates a percentage-of-total budget.
34    ///
35    /// # Errors
36    ///
37    /// Returns [`MemoryError::InvalidBudget`] unless `percent` is in `(0, 100]`.
38    pub fn percent_total(percent: f64) -> MemoryResult<Self> {
39        Ok(Self::PercentTotal(validate_percent(percent)? / 100.0))
40    }
41
42    /// Creates a percentage-of-available budget.
43    ///
44    /// # Errors
45    ///
46    /// Returns [`MemoryError::InvalidBudget`] unless `percent` is in `(0, 100]`.
47    pub fn percent_available(percent: f64) -> MemoryResult<Self> {
48        Ok(Self::PercentAvailable(validate_percent(percent)? / 100.0))
49    }
50
51    /// Resolves this request for a resource snapshot.
52    ///
53    /// # Errors
54    ///
55    /// Returns an error for zero budgets, invalid percentages, or unavailable
56    /// capacity telemetry.
57    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/// Host and optional accelerator budgets for one execution.
140#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
141pub struct MemoryPlan {
142    /// Host allocations, including source and staging buffers.
143    pub host: MemoryBudget,
144    /// Device allocations. Required only for accelerator execution.
145    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    /// Creates a host-only plan.
159    pub const fn host(host: MemoryBudget) -> Self {
160        Self { host, device: None }
161    }
162    /// Creates a host-and-device plan.
163    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}