use crate::lexer::Punctuator;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct Power(pub u8);
pub const LOWEST: Power = Power(0);
pub const OR: Power = Power(1);
pub const AND: Power = Power(2);
pub const NOT: Power = Power(3);
pub const COMPARISON: Power = Power(4);
pub const BITWISE: Power = Power(5);
pub const DISTANCE: Power = Power(5);
pub const ADDITIVE: Power = Power(6);
pub const MULTIPLICATIVE: Power = Power(7);
pub const CONCAT: Power = Power(8);
pub const COLLATE: Power = Power(9);
pub const UNARY: Power = Power(10);
pub fn infix_power(punctuator: Punctuator) -> Option<Power> {
let power = match punctuator {
Punctuator::Equal
| Punctuator::NotEqual
| Punctuator::Less
| Punctuator::LessEqual
| Punctuator::Greater
| Punctuator::GreaterEqual => COMPARISON,
Punctuator::BitAnd | Punctuator::BitOr | Punctuator::ShiftLeft | Punctuator::ShiftRight => {
BITWISE
}
Punctuator::L2Distance
| Punctuator::CosineDistance
| Punctuator::NegativeInnerProduct
| Punctuator::L1Distance
| Punctuator::HammingDistance
| Punctuator::JaccardDistance => DISTANCE,
Punctuator::Plus | Punctuator::Minus => ADDITIVE,
Punctuator::Star | Punctuator::Slash | Punctuator::Percent => MULTIPLICATIVE,
Punctuator::Concat | Punctuator::Arrow | Punctuator::DoubleArrow => CONCAT,
_ => return None,
};
Some(power)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_levels_are_in_the_published_order() {
let levels = [
LOWEST,
OR,
AND,
NOT,
COMPARISON,
BITWISE,
ADDITIVE,
MULTIPLICATIVE,
CONCAT,
COLLATE,
UNARY,
];
for pair in levels.windows(2) {
let (weaker, stronger) = (pair.first().copied(), pair.get(1).copied());
assert!(weaker < stronger, "{weaker:?} !< {stronger:?}");
}
}
#[test]
fn every_comparison_shares_one_level() {
for punctuator in [
Punctuator::Equal,
Punctuator::NotEqual,
Punctuator::Less,
Punctuator::LessEqual,
Punctuator::Greater,
Punctuator::GreaterEqual,
] {
assert_eq!(infix_power(punctuator), Some(COMPARISON), "{punctuator:?}");
}
}
#[test]
fn concatenation_binds_tighter_than_arithmetic() {
assert!(infix_power(Punctuator::Concat) > infix_power(Punctuator::Star));
assert!(infix_power(Punctuator::Star) > infix_power(Punctuator::Plus));
assert!(infix_power(Punctuator::Plus) > infix_power(Punctuator::BitOr));
}
#[test]
fn non_operators_have_no_power() {
for punctuator in [
Punctuator::LeftParen,
Punctuator::RightParen,
Punctuator::Comma,
Punctuator::Semicolon,
Punctuator::Dot,
Punctuator::BitNot,
] {
assert_eq!(infix_power(punctuator), None, "{punctuator:?}");
}
}
}