use std::sync::Arc;
use super::ast::{BinaryOp, Expression, UnaryOp};
use super::root::{DefaultExpressionRoot, ExpressionRoot};
use super::ParseError;
use crate::http::security::User;
#[derive(Debug, Clone)]
pub enum EvaluationError {
UnknownFunction(String),
}
impl std::fmt::Display for EvaluationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EvaluationError::UnknownFunction(name) => {
write!(f, "unknown function: '{}'", name)
}
}
}
}
impl std::error::Error for EvaluationError {}
pub struct ExpressionEvaluator {
root: Arc<dyn ExpressionRoot>,
}
impl ExpressionEvaluator {
pub fn new() -> Self {
ExpressionEvaluator {
root: Arc::new(DefaultExpressionRoot::new()),
}
}
pub fn with_root<R: ExpressionRoot + 'static>(root: R) -> Self {
ExpressionEvaluator {
root: Arc::new(root),
}
}
pub fn evaluate(
&self,
expr: &Expression,
user: Option<&User>,
) -> Result<bool, EvaluationError> {
match expr {
Expression::Boolean(value) => Ok(*value),
Expression::Function { name, args } => self
.root
.evaluate_function(name, args, user)
.ok_or_else(|| EvaluationError::UnknownFunction(name.clone())),
Expression::Binary { left, op, right } => {
let left_result = self.evaluate(left, user)?;
match op {
BinaryOp::And => {
if !left_result {
return Ok(false);
}
self.evaluate(right, user)
}
BinaryOp::Or => {
if left_result {
return Ok(true);
}
self.evaluate(right, user)
}
}
}
Expression::Unary { op, expr } => {
let result = self.evaluate(expr, user)?;
match op {
UnaryOp::Not => Ok(!result),
}
}
Expression::Group(inner) => self.evaluate(inner, user),
}
}
pub fn evaluate_str(&self, expr: &str, user: Option<&User>) -> Result<bool, ExpressionError> {
let parsed = super::SecurityExpression::parse(expr)?;
self.evaluate(parsed.ast(), user)
.map_err(ExpressionError::Evaluation)
}
}
impl Default for ExpressionEvaluator {
fn default() -> Self {
Self::new()
}
}
impl Clone for ExpressionEvaluator {
fn clone(&self) -> Self {
ExpressionEvaluator {
root: Arc::clone(&self.root),
}
}
}
#[derive(Debug)]
pub enum ExpressionError {
Parse(ParseError),
Evaluation(EvaluationError),
}
impl std::fmt::Display for ExpressionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ExpressionError::Parse(e) => write!(f, "parse error: {}", e),
ExpressionError::Evaluation(e) => write!(f, "evaluation error: {}", e),
}
}
}
impl std::error::Error for ExpressionError {}
impl From<ParseError> for ExpressionError {
fn from(err: ParseError) -> Self {
ExpressionError::Parse(err)
}
}
impl From<EvaluationError> for ExpressionError {
fn from(err: EvaluationError) -> Self {
ExpressionError::Evaluation(err)
}
}