Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 76 additions & 16 deletions ext/or-tools/routing.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,72 @@ namespace Rice::detail {
};
} // namespace Rice::detail

template<typename T>
operations_research::Constraint *make_constraint(operations_research::Solver &solver, operations_research::IntExpr *left, T right, const std::string &op) {
if (op == "==") {
return solver.MakeEquality(left, right);
} else if (op == "!=") {
return solver.MakeNonEquality(left, right);
} else if (op == "<") {
return solver.MakeLess(left, right);
} else if (op == "<=") {
return solver.MakeLessOrEqual(left, right);
} else if (op == ">") {
return solver.MakeGreater(left, right);
} else if (op == ">=") {
return solver.MakeGreaterOrEqual(left, right);
} else {
throw std::runtime_error{"Unknown operator"};
}
}

operations_research::IntExpr *make_int_expr(operations_research::Solver &solver, Object expression) {
if (expression.class_of().name() == "ORTools::Constant") {
return solver.MakeIntConst(Rice::detail::From_Ruby<int64_t>().convert(expression.call("value")));
} else if (expression.class_of().name() == "ORTools::Product") {
// MakeProd handles the left and right Bound() cases specially itself, so we don't need to optimize
// IntVar*Constant or Constant*IntVar ourselves
return solver.MakeProd(make_int_expr(solver, expression.call("left")), make_int_expr(solver, expression.call("right")));
} else if (expression.class_of().name() == "ORTools::Expression") {
const Array parts(expression.call("parts"));

switch (parts.size()) {
case 1:
// Unwrap unary Expression
return make_int_expr(solver, parts[0]);

case 2:
// ExpressionMethods#-(other) is implemented using + -(other), which we extract here mainly to help the common IntVar-IntVar case
if (parts[1].class_of().name() == "ORTools::Product" &&
parts[1].call("left").class_of().name() == "ORTools::Constant" &&
Rice::detail::From_Ruby<int64_t>().convert(parts[1].call("left").call("value")) == -1) {
return solver.MakeDifference(make_int_expr(solver, parts[0]), make_int_expr(solver, parts[1].call("right")));
} else {
return solver.MakeSum(make_int_expr(solver, parts[0]), make_int_expr(solver, parts[1]));
}

default:
// There's a sum(IntVar[]), but not a sum(IntExpr[])
if (std::all_of(parts.begin(), parts.end(), [](Object object) {
return Rice::detail::From_Ruby<operations_research::IntVar*>().is_convertible(object);
})) {
std::vector<operations_research::IntVar*> vars;
std::transform(parts.begin(), parts.end(), std::back_inserter(vars), [](Object object) {
return Rice::detail::From_Ruby<operations_research::IntVar*>().convert(object);
});
return solver.MakeSum(vars);
} else {
operations_research::IntExpr *curr = make_int_expr(solver, parts[0]);
for (size_t index = 1; index < parts.size(); index++) curr = solver.MakeSum(curr, make_int_expr(solver, parts[index]));
return curr;
}
}
} else {
// Anything else should be a Variable, ie. a RoutingIntVar
return Rice::detail::From_Ruby<operations_research::IntVar*>().convert(expression);
}
}

void init_routing(Rice::Module& m) {
auto rb_cRoutingSearchParameters = Rice::define_class_under<RoutingSearchParameters>(m, "RoutingSearchParameters");
auto rb_cIntVar = Rice::define_class_under<operations_research::IntVar>(m, "RoutingIntVar");
Expand Down Expand Up @@ -293,24 +359,18 @@ void init_routing(Rice::Module& m) {

Rice::define_class_under<operations_research::Solver>(m, "RoutingSolver")
.define_method(
"add",
[](operations_research::Solver& self, Object o) {
operations_research::Constraint* constraint;
if (o.respond_to("left")) {
operations_research::IntExpr* left(Rice::detail::From_Ruby<operations_research::IntVar*>().convert(o.call("left")));
operations_research::IntExpr* right(Rice::detail::From_Ruby<operations_research::IntVar*>().convert(o.call("right")));
std::string op = o.call("op").to_s().str();
if (op == "==") {
constraint = self.MakeEquality(left, right);
} else if (op == "<=") {
constraint = self.MakeLessOrEqual(left, right);
} else {
throw std::runtime_error{"Unknown operator"};
}
"add_constraint",
[](operations_research::Solver& self, operations_research::Constraint& constraint) {
self.AddConstraint(&constraint);
})
.define_method(
"_make_constraint",
[](operations_research::Solver& self, Object left, Object right, Symbol op) {
if (right.class_of().name() == "ORTools::Constant") {
return make_constraint(self, make_int_expr(self, left), Rice::detail::From_Ruby<int64_t>().convert(right.call("value")), op.str());
} else {
constraint = Rice::detail::From_Ruby<operations_research::Constraint*>().convert(o);
return make_constraint(self, make_int_expr(self, left), make_int_expr(self, right), op.str());
}
self.AddConstraint(constraint);
})
.define_method(
"fixed_duration_interval_var",
Expand Down
1 change: 1 addition & 0 deletions lib/or-tools.rb
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
# routing
require_relative "or_tools/routing_index_manager"
require_relative "or_tools/routing_model"
require_relative "or_tools/routing_solver"

# higher level interfaces
require_relative "or_tools/basic_scheduler"
Expand Down
14 changes: 14 additions & 0 deletions lib/or_tools/routing_solver.rb
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
module ORTools
class RoutingSolver
def add(comparison)
case comparison
when Comparison
add_constraint(_make_constraint(comparison.left, comparison.right, comparison.op))
when Constraint
add_constraint(comparison)
else
raise TypeError, "Not supported: RoutingSolver#add(#{comparison})"
end
end
end
end
202 changes: 202 additions & 0 deletions test/routing_constraints_test.rb
Original file line number Diff line number Diff line change
@@ -0,0 +1,202 @@
require_relative "test_helper"

class RoutingConstraintsTest < Minitest::Test
def test_no_extra_constraints
build_routing
solve

assert_equal :success, @routing.status
assert_equal [0, 2, 1, 0], route
end

def test_var_less_than_or_equal_var
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) <= @distance_dimension.cumul_var(@manager.node_to_index(2)))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_less_than_or_equal_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) <= 2455)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_equal_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) == 2451)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_equal_const_failure
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) == 2455)
solve

assert_equal :fail, @routing.status
end

def test_var_not_equal_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) != 2451)
solve

assert_equal :success, @routing.status
assert_equal [0, 2, 1, 0], route
end

def test_var_less_than_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) < 2455)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_greater_than_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) > 2455)
solve

assert_equal :success, @routing.status
assert_equal [0, 2, 1, 0], route
end

def test_var_greater_than_or_equal_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) >= 2455)
solve

assert_equal :success, @routing.status
assert_equal [0, 2, 1, 0], route
end

def test_var_plus_var
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) + @distance_dimension.cumul_var(@manager.node_to_index(2)) == (2451 + (2451 + 1745)))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_sum_vars
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(0)) + @distance_dimension.cumul_var(@manager.node_to_index(1)) + @distance_dimension.cumul_var(@manager.node_to_index(2)) == (2451 + (2451 + 1745)))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_plus_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) + 10000 == (2451 + 10000))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_plus_var
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) + @distance_dimension.cumul_var(@manager.node_to_index(2)) + 10000 == (2451 + (2451 + 1745) + 10000))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_minus_var
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(2)) - @distance_dimension.cumul_var(@manager.node_to_index(1)) == ((2451 + 1745) - 2451))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_minus_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1)) - 1000 == (2451 - 1000))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_product_const
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1))*2 <= 4910)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_const_product_var
build_routing
@routing.solver.add(2*@distance_dimension.cumul_var(@manager.node_to_index(1)) <= 4910)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_var_product_var
build_routing
@routing.solver.add(@distance_dimension.cumul_var(@manager.node_to_index(1))*@distance_dimension.cumul_var(@manager.node_to_index(2)) == (2451*(2451 + 1745)))
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

def test_negate_var
build_routing
@routing.solver.add(-@distance_dimension.cumul_var(@manager.node_to_index(1)) > -2455)
solve

assert_equal :success, @routing.status
assert_equal [0, 1, 2, 0], route
end

private

def build_routing
@manager = ORTools::RoutingIndexManager.new(3, 1, 0)
@routing = ORTools::RoutingModel.new(@manager)
transit_callback_index = @routing.register_transit_matrix([
[0, 2451, 731],
[2451, 0, 1745],
[731, 1745, 0],
])
@routing.set_arc_cost_evaluator_of_all_vehicles(transit_callback_index)
@routing.add_dimension(transit_callback_index, 0, 10000, true, "Distance")
@distance_dimension = @routing.mutable_dimension("Distance")
end

def solve
@solution = @routing.solve(first_solution_strategy: :path_cheapest_arc)
end

def route
route = []
index = @routing.start(0)
while !@routing.end?(index)
route << @manager.index_to_node(index)
index = @solution.value(@routing.next_var(index))
end
route << @manager.index_to_node(index)
route
end
end