diff --git a/ext/or-tools/routing.cpp b/ext/or-tools/routing.cpp index c7b23e7..da1fa44 100644 --- a/ext/or-tools/routing.cpp +++ b/ext/or-tools/routing.cpp @@ -97,6 +97,25 @@ namespace Rice::detail { }; } // namespace Rice::detail +template +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"}; + } +} + void init_routing(Rice::Module& m) { auto rb_cRoutingSearchParameters = Rice::define_class_under(m, "RoutingSearchParameters"); auto rb_cIntVar = Rice::define_class_under(m, "RoutingIntVar"); @@ -293,24 +312,18 @@ void init_routing(Rice::Module& m) { Rice::define_class_under(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().convert(o.call("left"))); - operations_research::IntExpr* right(Rice::detail::From_Ruby().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, Rice::detail::From_Ruby().convert(left), Rice::detail::From_Ruby().convert(right.call("value")), op.str()); } else { - constraint = Rice::detail::From_Ruby().convert(o); + return make_constraint(self, Rice::detail::From_Ruby().convert(left), Rice::detail::From_Ruby().convert(right), op.str()); } - self.AddConstraint(constraint); }) .define_method( "fixed_duration_interval_var", diff --git a/lib/or-tools.rb b/lib/or-tools.rb index 4b9dac1..93c1380 100644 --- a/lib/or-tools.rb +++ b/lib/or-tools.rb @@ -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" diff --git a/lib/or_tools/routing_solver.rb b/lib/or_tools/routing_solver.rb new file mode 100644 index 0000000..618c7f0 --- /dev/null +++ b/lib/or_tools/routing_solver.rb @@ -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 diff --git a/test/routing_constraints_test.rb b/test/routing_constraints_test.rb new file mode 100644 index 0000000..bd3791c --- /dev/null +++ b/test/routing_constraints_test.rb @@ -0,0 +1,112 @@ +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 + + 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