blob: 443e4e20a190db77a75c6dd1c61d8d4524dc69dd [file]
#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