Skip to main content

launchbound_space/
constraint.rs

1//! A deliberately small constraint language over integer dimensions.
2//!
3//! Grammar (no parentheses, two precedence levels):
4//!   constraint := arith CMP arith
5//!   arith      := term (('+' | '-') term)*
6//!   term       := atom (('*' | '/' | '%') atom)*
7//!   atom       := integer | dimension-name
8//!   CMP        := '==' | '!=' | '<=' | '>=' | '<' | '>'
9//!
10//! Evaluation is checked u64 arithmetic (division truncates); overflow or
11//! division/modulo by zero makes the constraint an error, never silently
12//! true or false.
13
14use 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    // First pass: fold * and % into a term list separated by +/-.
171    let mut terms: Vec<(char, u64)> = Vec::new(); // (sign-op, value)
172    let mut pending_op: Option<char> = None; // within-term * or %
173    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
247/// Evaluate a comparison-free arithmetic expression against a candidate's
248/// dimensions plus extra named variables (bench plans use this for grid
249/// shapes and buffer sizes, e.g. `elements / block_x`).
250pub 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}