| #include "tlbmc/expression/expression.h" |
| |
| #include <algorithm> |
| #include <cctype> |
| #include <cmath> |
| #include <cstddef> |
| #include <memory> |
| #include <string> |
| #include <system_error> // NOLINT: system_error is commonly used in BMC |
| #include <utility> |
| #include <vector> |
| |
| #include "absl/container/flat_hash_map.h" |
| #include "absl/container/flat_hash_set.h" |
| #include "absl/status/status.h" |
| #include "absl/status/statusor.h" |
| #include "absl/strings/charconv.h" |
| #include "absl/strings/str_cat.h" |
| #include "absl/strings/string_view.h" |
| #include "g3/macros.h" |
| |
| namespace milotic_tlbmc { |
| namespace expression::internal { |
| |
| // Represents a constant value. |
| class Constant final : public Expression { |
| public: |
| explicit Constant(double value) : value_(value) {} |
| absl::StatusOr<double> Evaluate( |
| const absl::flat_hash_map<std::string, double>& /*variable_maps*/) |
| const override { |
| return value_; |
| } |
| |
| void GetRequiredVariables( |
| absl::flat_hash_set<std::string>& /*variable_names*/) const override {} |
| |
| private: |
| const double value_; |
| }; |
| |
| // Represents a variable. |
| class Variable final : public Expression { |
| public: |
| explicit Variable(absl::string_view name) : name_(name) {} |
| absl::StatusOr<double> Evaluate( |
| const absl::flat_hash_map<std::string, double>& variable_maps) |
| const override { |
| if (auto it = variable_maps.find(name_); it != variable_maps.end()) { |
| return it->second; |
| } |
| return absl::NotFoundError(absl::StrCat("Variable not found: ", name_)); |
| } |
| |
| void GetRequiredVariables( |
| absl::flat_hash_set<std::string>& variable_names) const override { |
| variable_names.insert(name_); |
| } |
| |
| private: |
| const std::string name_; |
| }; |
| |
| // Represents a binary operation. |
| class BinaryOperation final : public Expression { |
| public: |
| enum class Operator { |
| kAdd, |
| kSubtract, |
| kMultiply, |
| kDivide, |
| kPower, |
| kEquals, |
| kGreaterThan, |
| kGreaterThanOrEqual, |
| kLessThan, |
| kLessThanOrEqual |
| }; |
| |
| BinaryOperation(Operator op, std::unique_ptr<Expression> lhs, |
| std::unique_ptr<Expression> rhs) |
| : operator_(op), lhs_(std::move(lhs)), rhs_(std::move(rhs)) {} |
| |
| absl::StatusOr<double> Evaluate( |
| const absl::flat_hash_map<std::string, double>& variable_maps) |
| const override { |
| ECCLESIA_ASSIGN_OR_RETURN(double lhs_val, lhs_->Evaluate(variable_maps)); |
| ECCLESIA_ASSIGN_OR_RETURN(double rhs_val, rhs_->Evaluate(variable_maps)); |
| |
| switch (operator_) { |
| case Operator::kAdd: |
| return lhs_val + rhs_val; |
| case Operator::kSubtract: |
| return lhs_val - rhs_val; |
| case Operator::kMultiply: |
| return lhs_val * rhs_val; |
| case Operator::kDivide: |
| if (rhs_val == 0) { |
| return absl::InvalidArgumentError("Division by zero."); |
| } |
| return lhs_val / rhs_val; |
| case Operator::kPower: |
| return std::pow(lhs_val, rhs_val); |
| case Operator::kEquals: |
| return lhs_val == rhs_val; |
| case Operator::kGreaterThan: |
| return lhs_val > rhs_val; |
| case Operator::kGreaterThanOrEqual: |
| return lhs_val >= rhs_val; |
| case Operator::kLessThan: |
| return lhs_val < rhs_val; |
| case Operator::kLessThanOrEqual: |
| return lhs_val <= rhs_val; |
| } |
| return absl::InternalError("Unreachable."); |
| } |
| |
| void GetRequiredVariables( |
| absl::flat_hash_set<std::string>& variable_names) const override { |
| lhs_->GetRequiredVariables(variable_names); |
| rhs_->GetRequiredVariables(variable_names); |
| } |
| |
| private: |
| const Operator operator_; |
| const std::unique_ptr<Expression> lhs_; |
| const std::unique_ptr<Expression> rhs_; |
| }; |
| |
| // Represents the Maximum function. |
| class Max final : public Expression { |
| public: |
| explicit Max(std::vector<std::unique_ptr<Expression>> operands) |
| : operands_(std::move(operands)) {} |
| |
| absl::StatusOr<double> Evaluate( |
| const absl::flat_hash_map<std::string, double>& variable_maps) |
| const override { |
| // operands are guaranteed to be non-empty by the parser. |
| ECCLESIA_ASSIGN_OR_RETURN(double max_val, |
| operands_[0]->Evaluate(variable_maps)); |
| for (const auto& operand : operands_) { |
| ECCLESIA_ASSIGN_OR_RETURN(double val, operand->Evaluate(variable_maps)); |
| max_val = std::max(max_val, val); |
| } |
| return max_val; |
| } |
| |
| void GetRequiredVariables( |
| absl::flat_hash_set<std::string>& variable_names) const override { |
| for (const auto& operand : operands_) { |
| operand->GetRequiredVariables(variable_names); |
| } |
| } |
| |
| private: |
| const std::vector<std::unique_ptr<Expression>> operands_; |
| }; |
| |
| // Represents the Condition function. |
| class Condition final : public Expression { |
| public: |
| explicit Condition(std::vector<std::unique_ptr<Expression>> operands) |
| : operands_(std::move(operands)) {} |
| |
| absl::StatusOr<double> Evaluate( |
| const absl::flat_hash_map<std::string, double>& variable_maps) |
| const override { |
| // operands are guaranteed to be non-empty by the parser. |
| if (operands_.size() != 3) { |
| return absl::InvalidArgumentError( |
| "Condition requires exactly 3 arguments."); |
| } |
| |
| ECCLESIA_ASSIGN_OR_RETURN(bool bool_val, |
| operands_[0]->Evaluate(variable_maps)); |
| if (bool_val) { |
| return operands_[1]->Evaluate(variable_maps); |
| } |
| return operands_[2]->Evaluate(variable_maps); |
| } |
| |
| void GetRequiredVariables( |
| absl::flat_hash_set<std::string>& variable_names) const override { |
| for (const auto& operand : operands_) { |
| operand->GetRequiredVariables(variable_names); |
| } |
| } |
| |
| private: |
| const std::vector<std::unique_ptr<Expression>> operands_; |
| }; |
| |
| // Consumes whitespace until the first non-whitespace character. |
| absl::string_view ConsumeWhitespace(absl::string_view s) { |
| int i = 0; |
| while (!s.empty() && i < s.size() && std::isspace(s[i]) != 0) { |
| i++; |
| } |
| return s.substr(i); |
| } |
| |
| // Returns the first character in the string, or '\0' if the string is empty. |
| char Peek(absl::string_view s) { |
| s = ConsumeWhitespace(s); |
| return s.empty() ? '\0' : s.front(); |
| } |
| |
| // Consumes the given character if it is present. Returns true if the character |
| // is found, false otherwise. |
| std::pair<absl::string_view, bool> Consume(absl::string_view s, char c) { |
| s = ConsumeWhitespace(s); |
| if (!s.empty() && s.front() == c) { |
| return std::make_pair(s.substr(1), true); |
| } |
| return std::make_pair(s, false); |
| } |
| |
| // Parses a primary, which is a number or a variable or an primary operator. |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParsePrimary(absl::string_view s) { |
| s = ConsumeWhitespace(s); |
| if (s.empty()) { |
| return absl::InvalidArgumentError("Unexpected end of expression."); |
| } |
| |
| // Parse a number. |
| if ((std::isdigit(s.front()) != 0) || s.front() == '.' || s.front() == '-') { |
| double value; |
| auto [ptr, ec] = absl::from_chars(s.data(), s.data() + s.size(), value); |
| if (ec == std::errc()) { |
| return std::make_pair(ptr, std::make_unique<Constant>(value)); |
| } |
| return absl::InvalidArgumentError( |
| absl::StrCat("Failed to parse number: ", ec)); |
| } |
| |
| // Parse a variable name or a operator name. |
| if (std::isalpha(s.front()) != 0) { |
| size_t len = 0; |
| while (len < s.size() && (std::isalnum(s[len]) != 0 || s[len] == '_')) { |
| len++; |
| } |
| absl::string_view name = s.substr(0, len); |
| s = s.substr(len); |
| if (name == "Maximum") { |
| std::pair<absl::string_view, bool> opening_parenthesis_found = |
| Consume(s, '('); |
| s = opening_parenthesis_found.first; |
| if (!opening_parenthesis_found.second) { |
| return absl::InvalidArgumentError("Expected '(' after Maximum."); |
| } |
| std::vector<std::unique_ptr<Expression>> args; |
| if (Peek(s) != ')') { |
| while (true) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> arg_pair; |
| ECCLESIA_ASSIGN_OR_RETURN(arg_pair, ParseExpression(s)); |
| s = arg_pair.first; |
| std::unique_ptr<Expression> arg = std::move(arg_pair.second); |
| args.push_back(std::move(arg)); |
| std::pair<absl::string_view, bool> comma_found = Consume(s, ','); |
| s = comma_found.first; |
| if (!comma_found.second) { |
| break; |
| } |
| } |
| } |
| std::pair<absl::string_view, bool> closing_parenthesis_found = |
| Consume(s, ')'); |
| s = closing_parenthesis_found.first; |
| if (!closing_parenthesis_found.second) { |
| return absl::InvalidArgumentError( |
| "Expected ')' after Maximum arguments."); |
| } |
| if (args.empty()) { |
| return absl::InvalidArgumentError( |
| "Maximum requires at least one argument."); |
| } |
| return std::make_pair(s, std::make_unique<Max>(std::move(args))); |
| } |
| |
| if (name == "Condition") { |
| std::pair<absl::string_view, bool> opening_parenthesis_found = |
| Consume(s, '('); |
| s = opening_parenthesis_found.first; |
| if (!opening_parenthesis_found.second) { |
| return absl::InvalidArgumentError("Expected '(' after Condition."); |
| } |
| std::vector<std::unique_ptr<Expression>> args; |
| if (Peek(s) != ')') { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> condition; |
| ECCLESIA_ASSIGN_OR_RETURN(condition, ParseConditionalOp(s)); |
| s = condition.first; |
| args.push_back(std::move(condition.second)); |
| |
| std::pair<absl::string_view, bool> comma_found = Consume(s, ','); |
| s = comma_found.first; |
| if (!comma_found.second) { |
| return absl::InvalidArgumentError( |
| "Expected ',' between Condition arguments."); |
| } |
| std::pair<absl::string_view, std::unique_ptr<Expression>> true_value; |
| ECCLESIA_ASSIGN_OR_RETURN(true_value, ParseExpression(s)); |
| s = true_value.first; |
| args.push_back(std::move(true_value.second)); |
| |
| std::pair<absl::string_view, bool> comma_found2 = Consume(s, ','); |
| s = comma_found2.first; |
| if (!comma_found2.second) { |
| return absl::InvalidArgumentError( |
| "Expected ',' between Condition arguments."); |
| } |
| std::pair<absl::string_view, std::unique_ptr<Expression>> false_value; |
| ECCLESIA_ASSIGN_OR_RETURN(false_value, ParseExpression(s)); |
| s = false_value.first; |
| args.push_back(std::move(false_value.second)); |
| } |
| std::pair<absl::string_view, bool> closing_parenthesis_found = |
| Consume(s, ')'); |
| s = closing_parenthesis_found.first; |
| if (!closing_parenthesis_found.second) { |
| return absl::InvalidArgumentError( |
| "Expected ')' after Condition arguments."); |
| } |
| if (args.size() != 3) { |
| return absl::InvalidArgumentError( |
| "Condition requires exactly 3 arguments."); |
| } |
| return std::make_pair(s, std::make_unique<Condition>(std::move(args))); |
| } |
| |
| return std::make_pair(s, std::make_unique<Variable>(name)); |
| } |
| |
| std::pair<absl::string_view, bool> opening_parenthesis_found = |
| Consume(s, '('); |
| s = opening_parenthesis_found.first; |
| if (opening_parenthesis_found.second) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> expr; |
| ECCLESIA_ASSIGN_OR_RETURN(expr, ParseExpression(s)); |
| s = expr.first; |
| std::pair<absl::string_view, bool> closing_parenthesis_found = |
| Consume(s, ')'); |
| s = closing_parenthesis_found.first; |
| if (!closing_parenthesis_found.second) { |
| return absl::InvalidArgumentError("Mismatched parentheses."); |
| } |
| return std::make_pair(s, std::move(expr.second)); |
| } |
| |
| return absl::InvalidArgumentError( |
| absl::StrCat("Unexpected token: ", s.substr(0, 1))); |
| } |
| |
| // Parses an exponent, which is a primary followed by a sequence of power |
| // operators. |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParseExponent(absl::string_view s) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> lhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(lhs_result, ParsePrimary(s)); |
| std::unique_ptr<Expression> lhs = std::move(lhs_result.second); |
| s = lhs_result.first; |
| |
| while (true) { |
| if (Peek(s) == '^') { |
| s = Consume(s, '^').first; |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParsePrimary(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>(BinaryOperation::Operator::kPower, |
| std::move(lhs), std::move(rhs)); |
| } else { |
| break; |
| } |
| } |
| return std::make_pair(s, std::move(lhs)); |
| } |
| |
| // Parses a factor, which is a primary followed by a sequence of multiplication |
| // and division operators. |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParseFactor(absl::string_view s) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> lhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(lhs_result, ParseExponent(s)); |
| std::unique_ptr<Expression> lhs = std::move(lhs_result.second); |
| s = lhs_result.first; |
| |
| while (true) { |
| if (Peek(s) == '*') { |
| s = Consume(s, '*').first; |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParseExponent(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>( |
| BinaryOperation::Operator::kMultiply, std::move(lhs), std::move(rhs)); |
| } else if (Peek(s) == '/') { |
| s = Consume(s, '/').first; |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParseExponent(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>( |
| BinaryOperation::Operator::kDivide, std::move(lhs), std::move(rhs)); |
| } else { |
| break; |
| } |
| } |
| return std::make_pair(s, std::move(lhs)); |
| } |
| |
| // Parses a term, which is a factor followed by a sequence of binary |
| // operations. |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParseTerm(absl::string_view s) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> lhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(lhs_result, ParseFactor(s)); |
| s = lhs_result.first; |
| std::unique_ptr<Expression> lhs = std::move(lhs_result.second); |
| |
| while (true) { |
| if (Peek(s) == '+') { |
| s = Consume(s, '+').first; |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParseFactor(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>(BinaryOperation::Operator::kAdd, |
| std::move(lhs), std::move(rhs)); |
| } else if (Peek(s) == '-') { |
| s = Consume(s, '-').first; |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParseFactor(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>( |
| BinaryOperation::Operator::kSubtract, std::move(lhs), std::move(rhs)); |
| } else { |
| break; |
| } |
| } |
| return std::make_pair(s, std::move(lhs)); |
| } |
| |
| // Parses a conditional operation which is a Term followed by a conditional |
| // operation. |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParseConditionalOp(absl::string_view s) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> lhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(lhs_result, ParseTerm(s)); |
| s = lhs_result.first; |
| std::unique_ptr<Expression> lhs = std::move(lhs_result.second); |
| |
| BinaryOperation::Operator op = BinaryOperation::Operator::kGreaterThan; |
| if (Peek(s) == '>') { |
| s = Consume(s, '>').first; |
| if (Peek(s) == '=') { |
| s = Consume(s, '=').first; |
| op = BinaryOperation::Operator::kGreaterThanOrEqual; |
| } |
| } else if (Peek(s) == '<') { |
| s = Consume(s, '<').first; |
| op = BinaryOperation::Operator::kLessThan; |
| if (Peek(s) == '=') { |
| s = Consume(s, '=').first; |
| op = BinaryOperation::Operator::kLessThanOrEqual; |
| } |
| } else if (Peek(s) == '=') { |
| s = Consume(s, '=').first; |
| if (Peek(s) == '=') { |
| s = Consume(s, '=').first; |
| op = BinaryOperation::Operator::kEquals; |
| } else { |
| return absl::InvalidArgumentError(absl::StrCat( |
| "Expected comparison operator, '==', '>', '<', '>=', or '<='. Got =", |
| std::string(1, Peek(s)))); |
| } |
| } else { |
| return absl::InvalidArgumentError(absl::StrCat( |
| "Expected comparison operator, '==', '>', '<', '>=', or '<='. Got ", |
| std::string(1, Peek(s)))); |
| } |
| std::pair<absl::string_view, std::unique_ptr<Expression>> rhs_result; |
| ECCLESIA_ASSIGN_OR_RETURN(rhs_result, ParseTerm(s)); |
| s = rhs_result.first; |
| std::unique_ptr<Expression> rhs = std::move(rhs_result.second); |
| lhs = std::make_unique<BinaryOperation>(op, std::move(lhs), std::move(rhs)); |
| |
| return std::make_pair(s, std::move(lhs)); |
| } |
| |
| absl::StatusOr<std::pair<absl::string_view, std::unique_ptr<Expression>>> |
| ParseExpression(absl::string_view s) { |
| return ParseTerm(s); |
| } |
| |
| } // namespace expression::internal |
| |
| absl::StatusOr<std::unique_ptr<Expression>> Parse(absl::string_view expr) { |
| std::pair<absl::string_view, std::unique_ptr<Expression>> result; |
| ECCLESIA_ASSIGN_OR_RETURN(result, |
| expression::internal::ParseExpression(expr)); |
| |
| // Remove any trailing whitespace. |
| absl::string_view left_over = |
| expression::internal::ConsumeWhitespace(result.first); |
| if (!left_over.empty()) { |
| return absl::InvalidArgumentError( |
| absl::StrCat("Unexpected trailing characters: ", left_over)); |
| } |
| return std::move(result.second); |
| } |
| |
| } // namespace milotic_tlbmc |