1use crate::spec::Value;
15use crate::{Config, SpaceError};
16
17#[derive(Debug, Clone, PartialEq, Eq)]
18enum Token {
19 Num(u64),
20 Ident(String),
21 Op(char),
22 Cmp(&'static str),
23}
24
25#[derive(Debug, Clone)]
26pub struct Constraint {
27 text: String,
28 lhs: Vec<Token>,
29 cmp: &'static str,
30 rhs: Vec<Token>,
31}
32
33impl Constraint {
34 pub fn text(&self) -> &str {
35 &self.text
36 }
37
38 pub fn parse(expr: &str, dims: &[&str]) -> Result<Self, SpaceError> {
39 let err = |reason: &str| SpaceError::Constraint {
40 expr: expr.to_string(),
41 reason: reason.to_string(),
42 };
43 let tokens = tokenize(expr).map_err(|r| err(&r))?;
44 let cmp_pos = tokens
45 .iter()
46 .position(|t| matches!(t, Token::Cmp(_)))
47 .ok_or_else(|| err("no comparison operator"))?;
48 let Token::Cmp(cmp) = tokens[cmp_pos] else {
49 unreachable!()
50 };
51 if tokens.iter().filter(|t| matches!(t, Token::Cmp(_))).count() != 1 {
52 return Err(err("exactly one comparison operator required"));
53 }
54 let lhs = tokens[..cmp_pos].to_vec();
55 let rhs = tokens[cmp_pos + 1..].to_vec();
56 for side in [&lhs, &rhs] {
57 if side.is_empty() {
58 return Err(err("empty side of comparison"));
59 }
60 for t in side {
61 if let Token::Ident(name) = t
62 && !dims.contains(&name.as_str())
63 {
64 return Err(err(&format!("unknown dimension `{name}`")));
65 }
66 }
67 }
68 Ok(Constraint {
69 text: expr.to_string(),
70 lhs,
71 cmp,
72 rhs,
73 })
74 }
75
76 pub fn eval(&self, config: &Config) -> Result<bool, SpaceError> {
77 let resolve = config_resolver(config);
78 let l = eval_arith(&self.lhs, &resolve, &self.text)?;
79 let r = eval_arith(&self.rhs, &resolve, &self.text)?;
80 Ok(match self.cmp {
81 "==" => l == r,
82 "!=" => l != r,
83 "<=" => l <= r,
84 ">=" => l >= r,
85 "<" => l < r,
86 ">" => l > r,
87 _ => unreachable!(),
88 })
89 }
90}
91
92fn tokenize(expr: &str) -> Result<Vec<Token>, String> {
93 let mut tokens = Vec::new();
94 let bytes = expr.as_bytes();
95 let mut i = 0;
96 while i < bytes.len() {
97 let c = bytes[i] as char;
98 match c {
99 ' ' | '\t' => i += 1,
100 '0'..='9' => {
101 let start = i;
102 while i < bytes.len() && bytes[i].is_ascii_digit() {
103 i += 1;
104 }
105 let n: u64 = expr[start..i]
106 .parse()
107 .map_err(|_| "integer literal too large".to_string())?;
108 tokens.push(Token::Num(n));
109 }
110 'a'..='z' | '_' => {
111 let start = i;
112 while i < bytes.len()
113 && (bytes[i].is_ascii_lowercase()
114 || bytes[i].is_ascii_digit()
115 || bytes[i] == b'_')
116 {
117 i += 1;
118 }
119 tokens.push(Token::Ident(expr[start..i].to_string()));
120 }
121 '*' | '/' | '%' | '+' | '-' => {
122 tokens.push(Token::Op(c));
123 i += 1;
124 }
125 '=' | '!' | '<' | '>' => {
126 let two = &expr[i..(i + 2).min(expr.len())];
127 let cmp = match two {
128 "==" => Some("=="),
129 "!=" => Some("!="),
130 "<=" => Some("<="),
131 ">=" => Some(">="),
132 _ => None,
133 };
134 if let Some(cmp) = cmp {
135 tokens.push(Token::Cmp(cmp));
136 i += 2;
137 } else if c == '<' {
138 tokens.push(Token::Cmp("<"));
139 i += 1;
140 } else if c == '>' {
141 tokens.push(Token::Cmp(">"));
142 i += 1;
143 } else {
144 return Err(format!("unexpected character `{c}`"));
145 }
146 }
147 other => return Err(format!("unexpected character `{other}`")),
148 }
149 }
150 Ok(tokens)
151}
152
153fn eval_arith(
154 tokens: &[Token],
155 resolve: &dyn Fn(&str) -> Result<u64, String>,
156 text: &str,
157) -> Result<u64, SpaceError> {
158 let err = |reason: String| SpaceError::Constraint {
159 expr: text.to_string(),
160 reason,
161 };
162 let atom = |t: &Token| -> Result<u64, SpaceError> {
163 match t {
164 Token::Num(n) => Ok(*n),
165 Token::Ident(name) => resolve(name).map_err(err),
166 Token::Op(_) | Token::Cmp(_) => Err(err("misplaced operator".into())),
167 }
168 };
169
170 let mut terms: Vec<(char, u64)> = Vec::new(); let mut pending_op: Option<char> = None; let mut sign: char = '+';
174 let mut current: Option<u64> = None;
175 for t in tokens {
176 match t {
177 Token::Op(op @ ('*' | '/' | '%')) => {
178 if current.is_none() {
179 return Err(err(format!("`{op}` with no left operand")));
180 }
181 pending_op = Some(*op);
182 }
183 Token::Op(op @ ('+' | '-')) => {
184 let value = current
185 .take()
186 .ok_or_else(|| err(format!("`{op}` with no left operand")))?;
187 terms.push((sign, value));
188 sign = *op;
189 pending_op = None;
190 }
191 atom_token => {
192 let v = atom(atom_token)?;
193 current = Some(match (current, pending_op.take()) {
194 (None, None) => v,
195 (Some(acc), Some('*')) => acc
196 .checked_mul(v)
197 .ok_or_else(|| err("multiplication overflow".into()))?,
198 (Some(acc), Some('%')) => {
199 if v == 0 {
200 return Err(err("modulo by zero".into()));
201 }
202 acc % v
203 }
204 (Some(acc), Some('/')) => {
205 if v == 0 {
206 return Err(err("division by zero".into()));
207 }
208 acc / v
209 }
210 (Some(_), None) => {
211 return Err(err("two operands with no operator".into()));
212 }
213 (None, Some(_)) => unreachable!(),
214 (Some(_), Some(_)) => unreachable!(),
215 });
216 }
217 }
218 }
219 let value = current.ok_or_else(|| err("trailing operator".into()))?;
220 terms.push((sign, value));
221
222 let mut acc: u64 = 0;
223 for (op, v) in terms {
224 acc = match op {
225 '+' => acc
226 .checked_add(v)
227 .ok_or_else(|| err("addition overflow".into()))?,
228 '-' => acc
229 .checked_sub(v)
230 .ok_or_else(|| err("subtraction underflow".into()))?,
231 _ => unreachable!(),
232 };
233 }
234 Ok(acc)
235}
236
237fn config_resolver(config: &Config) -> impl Fn(&str) -> Result<u64, String> + '_ {
238 move |name: &str| match config.get(name) {
239 Some(Value::Int(n)) => Ok(*n),
240 Some(Value::Str(_)) => Err(format!(
241 "dimension `{name}` is a string and cannot be used in arithmetic"
242 )),
243 None => Err(format!("dimension `{name}` missing from config")),
244 }
245}
246
247pub fn eval_arith_expr(
251 expr: &str,
252 config: &Config,
253 extra: &std::collections::BTreeMap<String, u64>,
254) -> Result<u64, SpaceError> {
255 let err = |reason: &str| SpaceError::Constraint {
256 expr: expr.to_string(),
257 reason: reason.to_string(),
258 };
259 let tokens = tokenize(expr).map_err(|r| err(&r))?;
260 if tokens.iter().any(|t| matches!(t, Token::Cmp(_))) {
261 return Err(err("comparison operators are not allowed here"));
262 }
263 let base = config_resolver(config);
264 let resolve = move |name: &str| -> Result<u64, String> {
265 if let Some(v) = extra.get(name) {
266 return Ok(*v);
267 }
268 base(name)
269 };
270 eval_arith(&tokens, &resolve, expr)
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276 use crate::KernelSpec;
277
278 fn config_with(block: u64, tile: u64) -> Config {
279 let spec = KernelSpec::from_toml_str(
280 "t",
281 &format!(
282 r#"
283 [kernel]
284 name = "t"
285 entry = "t"
286 domain = 1
287 [dims.block_x]
288 values = [{block}]
289 [dims.tile]
290 values = [{tile}]
291 "#
292 ),
293 )
294 .unwrap();
295 crate::enumerate(&spec).unwrap().into_iter().next().unwrap()
296 }
297
298 #[test]
299 fn arithmetic_and_comparisons() {
300 let dims = ["block_x", "tile"];
301 let c = config_with(64, 256);
302 for (expr, expected) in [
303 ("tile % block_x == 0", true),
304 ("tile % block_x != 0", false),
305 ("block_x * tile <= 16384", true),
306 ("block_x * tile < 16384", false),
307 ("tile - block_x == 192", true),
308 ("tile + block_x >= 320", true),
309 ("block_x > 32", true),
310 ] {
311 let parsed = Constraint::parse(expr, &dims).unwrap();
312 assert_eq!(parsed.eval(&c).unwrap(), expected, "{expr}");
313 }
314 }
315
316 #[test]
317 fn rejects_unknown_dimension_and_junk() {
318 let dims = ["block_x"];
319 assert!(Constraint::parse("bogus == 1", &dims).is_err());
320 assert!(Constraint::parse("block_x == ", &dims).is_err());
321 assert!(Constraint::parse("block_x", &dims).is_err());
322 assert!(Constraint::parse("block_x == 1 == 2", &dims).is_err());
323 assert!(Constraint::parse("block_x @ 2", &dims).is_err());
324 }
325
326 #[test]
327 fn division_by_zero_is_an_error_not_a_verdict() {
328 let dims = ["block_x", "tile"];
329 let c = config_with(64, 0);
330 let parsed = Constraint::parse("block_x % tile == 0", &dims).unwrap();
331 assert!(parsed.eval(&c).is_err());
332 }
333}