1use std::fmt;
4use std::str::FromStr;
5
6const KI: u64 = 1024;
7const MI: u64 = KI * 1024;
8const GI: u64 = MI * 1024;
9const TI: u64 = GI * 1024;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub struct MemoryBudget(u64);
19
20const BYTES_PER_SCORE: u64 = 8;
22
23const MIN_CHUNK_SCORES: usize = 4096;
26
27impl MemoryBudget {
28 pub fn bytes(&self) -> u64 {
29 self.0
30 }
31
32 pub fn chunk_scores(&self, threads: usize, cap: usize) -> usize {
39 let threads = threads.max(1) as u64;
40 let per_thread = self.0 / (threads * BYTES_PER_SCORE);
41 (per_thread as usize).clamp(MIN_CHUNK_SCORES, cap)
42 }
43}
44
45const MIN_WINDOW_BASES: usize = 1 << 20;
48
49impl MemoryBudget {
50 pub fn window_bases(&self, cap: usize) -> usize {
56 let share = (self.0 / 4) as usize;
59 share.clamp(MIN_WINDOW_BASES, cap)
60 }
61}
62
63impl fmt::Display for MemoryBudget {
64 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65 let b = self.0;
66 if b >= GI && b.is_multiple_of(GI) {
67 write!(f, "{}G", b / GI)
68 } else if b >= MI && b.is_multiple_of(MI) {
69 write!(f, "{}M", b / MI)
70 } else if b >= KI && b.is_multiple_of(KI) {
71 write!(f, "{}K", b / KI)
72 } else {
73 write!(f, "{b}")
74 }
75 }
76}
77
78#[derive(Debug, PartialEq, Eq)]
80pub struct ParseBudgetError(String);
81
82impl fmt::Display for ParseBudgetError {
83 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
84 write!(
85 f,
86 "{}: expected a size like 8G, 512M, 64K or a plain byte count",
87 self.0
88 )
89 }
90}
91
92impl std::error::Error for ParseBudgetError {}
93
94impl FromStr for MemoryBudget {
95 type Err = ParseBudgetError;
96
97 fn from_str(s: &str) -> Result<Self, Self::Err> {
103 let err = || ParseBudgetError(s.to_string());
104 let t = s.trim();
105 if t.is_empty() {
106 return Err(err());
107 }
108 let t = t.strip_suffix(['b', 'B']).unwrap_or(t);
110 let t = t.strip_suffix(['i', 'I']).unwrap_or(t);
111
112 let (digits, multiplier) = match t.chars().last().ok_or_else(err)? {
113 'k' | 'K' => (&t[..t.len() - 1], KI),
114 'm' | 'M' => (&t[..t.len() - 1], MI),
115 'g' | 'G' => (&t[..t.len() - 1], GI),
116 't' | 'T' => (&t[..t.len() - 1], TI),
117 _ => (t, 1),
118 };
119
120 let value: u64 = digits.trim().parse().map_err(|_| err())?;
121 let bytes = value.checked_mul(multiplier).ok_or_else(err)?;
122 if bytes == 0 {
123 return Err(err());
124 }
125 Ok(MemoryBudget(bytes))
126 }
127}
128
129#[cfg(test)]
130mod tests {
131 use super::*;
132
133 #[test]
134 fn test_parses_suffixes() {
135 let cases = [
136 ("8G", 8 * GI),
137 ("8g", 8 * GI),
138 ("8GB", 8 * GI),
139 ("8GiB", 8 * GI),
140 ("512M", 512 * MI),
141 ("64K", 64 * KI),
142 ("2T", 2 * TI),
143 ("1048576", 1048576),
144 (" 4G ", 4 * GI),
145 ];
146 for (text, expected) in cases {
147 assert_eq!(
148 text.parse::<MemoryBudget>().map(|b| b.bytes()),
149 Ok(expected),
150 "for {text:?}"
151 );
152 }
153 }
154
155 #[test]
156 fn test_rejects_nonsense() {
157 for text in ["", " ", "G", "8X", "-1", "1.5G", "eight", "0", "0G"] {
158 assert!(
159 text.parse::<MemoryBudget>().is_err(),
160 "{text:?} should not parse"
161 );
162 }
163 }
164
165 #[test]
166 fn test_rejects_overflow_rather_than_wrapping() {
167 assert!("99999999999T".parse::<MemoryBudget>().is_err());
168 }
169
170 #[test]
171 fn test_display_round_trips() {
172 for text in ["8G", "512M", "64K", "1023"] {
173 let parsed: MemoryBudget = text.parse().unwrap();
174 assert_eq!(parsed.to_string(), text);
175 }
176 }
177
178 #[test]
179 fn test_chunk_scores_scales_with_budget_and_threads() {
180 let cap = 1 << 20;
181 let eight_g: MemoryBudget = "8G".parse().unwrap();
182 assert_eq!(eight_g.chunk_scores(10, cap), cap);
184
185 let small: MemoryBudget = "64M".parse().unwrap();
188 assert_eq!(small.chunk_scores(8, cap), (64 * MI / (8 * 8)) as usize);
189
190 let tiny: MemoryBudget = "1K".parse().unwrap();
192 assert_eq!(tiny.chunk_scores(16, cap), MIN_CHUNK_SCORES);
193 }
194
195 #[test]
196 fn test_window_bases_scales_and_clamps() {
197 let cap = 64 << 20;
198 let big: MemoryBudget = "8G".parse().unwrap();
199 assert_eq!(big.window_bases(cap), cap, "a large budget hits the cap");
200
201 let mid: MemoryBudget = "64M".parse().unwrap();
202 assert_eq!(mid.window_bases(cap), (64 * MI / 4) as usize);
203
204 let tiny: MemoryBudget = "1K".parse().unwrap();
205 assert_eq!(
206 tiny.window_bases(cap),
207 MIN_WINDOW_BASES,
208 "floors rather than collapsing"
209 );
210 }
211
212 #[test]
213 fn test_chunk_scores_handles_zero_threads() {
214 let b: MemoryBudget = "8G".parse().unwrap();
215 assert_eq!(b.chunk_scores(0, 1 << 20), 1 << 20);
216 }
217}