From 95d806dd96e8eb38eb8549a654a8df3243ec4835 Mon Sep 17 00:00:00 2001 From: Vaios Papaspyros <8146703+bpapaspyros@users.noreply.github.com> Date: Mon, 24 Mar 2025 11:59:14 +0100 Subject: [PATCH] feat: update translators to include new trajectory msgs (#194) --- CHANGELOG.md | 6 ++ aica-package.toml | 2 +- source/modulo_core/CMakeLists.txt | 5 +- .../translators/message_readers.hpp | 17 +++++ .../translators/message_writers.hpp | 15 +++++ .../translators/message_readers.py | 22 +++++++ .../translators/message_writers.py | 21 +++++- source/modulo_core/package.xml | 1 + .../src/translators/message_readers.cpp | 37 +++++++++++ .../src/translators/message_writers.cpp | 29 ++++++++ .../test/cpp/translators/test_messages.cpp | 66 +++++++++++++++++++ source/modulo_core/test/python/conftest.py | 20 ++++++ .../test/python/translators/test_messages.py | 54 +++++++++++++++ 13 files changed, 292 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 975d4690e..4d6906a0f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,6 +33,7 @@ Release Versions: - feat(controllers)!: add TF broadcaster in BaseControllerInterface (#170) - test(controllers): add TF listener and broadcaster tests (#172) - fix(controllers)!: remove input validity setter (#209) +- feat(core): update translators to include new trajectory msgs (#194) ## 5.4.1 @@ -103,6 +104,11 @@ base class. - feat(components): verify return value of callbacks (#206) - fix(controllers): move input validity period to base (#207) +- feat(controllers)!: remove robot description parameter (#186) +- feat(controllers): add TF listener interface in BaseControllerInterface (#169) +- feat(controllers): add TF broadcaster in BaseControllerInterface (#170) +- test(controllers): add TF listener and broadcaster tests (#172) +- feat(controllers): use parent node for tf listener (#190) ## 5.2.0 diff --git a/aica-package.toml b/aica-package.toml index 788bdc931..6debf33ed 100644 --- a/aica-package.toml +++ b/aica-package.toml @@ -12,7 +12,7 @@ type = "ros" image = "v2.0.6-jazzy" [build.dependencies] -"@aica/foss/control-libraries" = "v9.0.0" +"@aica/foss/control-libraries" = "v10.0.0-rc0001" [build.packages.modulo_components] source = "./source/modulo_components" diff --git a/source/modulo_core/CMakeLists.txt b/source/modulo_core/CMakeLists.txt index 4d507b645..490640442 100644 --- a/source/modulo_core/CMakeLists.txt +++ b/source/modulo_core/CMakeLists.txt @@ -25,6 +25,7 @@ find_package(std_msgs REQUIRED) find_package(rclcpp REQUIRED) find_package(rclcpp_lifecycle REQUIRED) find_package(tf2_msgs REQUIRED) +find_package(trajectory_msgs REQUIRED) find_package(modulo_interfaces REQUIRED) find_package(control_libraries 9.0.0 REQUIRED COMPONENTS state_representation) @@ -56,6 +57,7 @@ ament_target_dependencies( sensor_msgs std_msgs tf2_msgs + trajectory_msgs modulo_interfaces ) @@ -83,7 +85,7 @@ if(BUILD_TESTING) ament_add_gtest(test_modulo_core ${TEST_CPP_SOURCES}) target_include_directories(test_modulo_core PRIVATE include) target_link_libraries(test_modulo_core ${PROJECT_NAME} clproto state_representation) - ament_target_dependencies(test_modulo_core geometry_msgs sensor_msgs std_msgs rclcpp rclcpp_lifecycle tf2_msgs) + ament_target_dependencies(test_modulo_core geometry_msgs sensor_msgs std_msgs rclcpp rclcpp_lifecycle tf2_msgs trajectory_msgs) # prevent pluginlib from using boost target_compile_definitions(test_modulo_core PUBLIC "PLUGINLIB__DISABLE_BOOST_FUNCTIONS") @@ -105,6 +107,7 @@ ament_export_dependencies( rclcpp rclcpp_lifecycle tf2_msgs + trajectory_msgs modulo_interfaces ) diff --git a/source/modulo_core/include/modulo_core/translators/message_readers.hpp b/source/modulo_core/include/modulo_core/translators/message_readers.hpp index 850509a14..949fc2baf 100644 --- a/source/modulo_core/include/modulo_core/translators/message_readers.hpp +++ b/source/modulo_core/include/modulo_core/translators/message_readers.hpp @@ -13,6 +13,8 @@ #include #include #include +#include +#include #include #include @@ -21,6 +23,8 @@ #include #include #include +#include +#include #include "modulo_core/EncodedState.hpp" #include "modulo_core/exceptions.hpp" @@ -104,6 +108,13 @@ void read_message(state_representation::CartesianState& state, const geometry_ms */ void read_message(state_representation::JointState& state, const sensor_msgs::msg::JointState& message); +/** + * @brief Convert a ROS trajectory_msgs::msg::JointTrajectory to a JointTrajectory + * @param state The JointTrajectory to populate + * @param message The ROS message to read from + */ +void read_message(state_representation::JointTrajectory& state, const trajectory_msgs::msg::JointTrajectory& message); + /** * @brief Template function to convert a ROS std_msgs::msg::T to a Parameter * @tparam T All types of parameters supported in ROS std messages @@ -436,6 +447,12 @@ inline void read_message(std::shared_ptr& state, co case StateType::JACOBIAN: safe_dynamic_cast(state, new_state); break; + case StateType::CARTESIAN_TRAJECTORY: + safe_dynamic_cast(state, new_state); + break; + case StateType::JOINT_TRAJECTORY: + safe_dynamic_cast(state, new_state); + break; case StateType::PARAMETER: { auto param_ptr = std::dynamic_pointer_cast(state); switch (param_ptr->get_parameter_type()) { diff --git a/source/modulo_core/include/modulo_core/translators/message_writers.hpp b/source/modulo_core/include/modulo_core/translators/message_writers.hpp index df44c55ce..5052b37c5 100644 --- a/source/modulo_core/include/modulo_core/translators/message_writers.hpp +++ b/source/modulo_core/include/modulo_core/translators/message_writers.hpp @@ -13,11 +13,15 @@ #include #include #include +#include +#include #include #include #include #include +#include +#include #include "modulo_core/EncodedState.hpp" #include "modulo_core/exceptions.hpp" @@ -150,6 +154,17 @@ void write_message( void write_message( tf2_msgs::msg::TFMessage& message, const state_representation::CartesianState& state, const rclcpp::Time& time); +/** + * @brief Convert a CartesianTrajectory to a ROS trajectory_msgs::msg::JointTrajectory + * @param message The ROS message to populate + * @param state The state to read from + * @param time The time of the message + * @throws modulo_core::exceptions::MessageTranslationException if the provided state is empty. + */ +void write_message( + trajectory_msgs::msg::JointTrajectory& message, const state_representation::JointTrajectory& state, + const rclcpp::Time& time); + /** * @brief Convert a Parameter to a ROS equivalent representation * @tparam T All types of parameters supported in ROS std messages diff --git a/source/modulo_core/modulo_core/translators/message_readers.py b/source/modulo_core/modulo_core/translators/message_readers.py index 30453b694..b340442a3 100644 --- a/source/modulo_core/modulo_core/translators/message_readers.py +++ b/source/modulo_core/modulo_core/translators/message_readers.py @@ -1,11 +1,13 @@ from typing import List, TypeVar, Union import clproto +import datetime import geometry_msgs.msg as geometry import state_representation as sr from modulo_core import EncodedState from modulo_core.exceptions import MessageTranslationError from sensor_msgs.msg import JointState +import trajectory_msgs.msg as trajectory DataT = TypeVar('DataT') MsgT = TypeVar('MsgT') @@ -74,6 +76,26 @@ def read_message(state: StateT, message: MsgT) -> StateT: state.set_torques(message.effort) except Exception as e: raise MessageTranslationError(f"{e}") + elif isinstance(message, trajectory.JointTrajectory) and isinstance(state, sr.JointTrajectory): + try: + state.set_joint_names(list(message.joint_names)) + time_from_start = 0 + for i, point in enumerate(message.points): + time_head = point.time_from_start.sec + point.time_from_start.nanosec * 1e-9 + duration = time_head - time_from_start + time_from_start = time_head + joint_state = sr.JointState(f"point_{i}", state.get_joint_names()) + joint_state.set_positions(point.positions) + joint_state.set_velocities(point.velocities) + joint_state.set_accelerations(point.accelerations) + joint_state.set_torques(point.effort) + state.add_point( + joint_state, + datetime.timedelta(seconds=duration) + ) + state.set_name(message.header.frame_id) + except Exception as e: + raise MessageTranslationError(f"{e}") else: raise MessageTranslationError("The provided combination of state type and message type is not supported") except MessageTranslationError: diff --git a/source/modulo_core/modulo_core/translators/message_writers.py b/source/modulo_core/modulo_core/translators/message_writers.py index 5cbe43253..555d3516c 100644 --- a/source/modulo_core/modulo_core/translators/message_writers.py +++ b/source/modulo_core/modulo_core/translators/message_writers.py @@ -8,6 +8,7 @@ from modulo_core import EncodedState from modulo_core.exceptions import MessageTranslationError from sensor_msgs.msg import JointState +import trajectory_msgs.msg as trajectory DataT = TypeVar('DataT') MsgT = TypeVar('MsgT') @@ -55,6 +56,10 @@ def get_clproto_msg_type(state: StateT) -> clproto.MessageType: return clproto.MessageType.DIGITAL_IO_STATE_MESSAGE elif state_type == sr.StateType.ANALOG_IO_STATE: return clproto.MessageType.ANALOG_IO_STATE_MESSAGE + elif state_type == sr.StateType.CARTESIAN_TRAJECTORY: + return clproto.MessageType.CARTESIAN_TRAJECTORY_MESSAGE + elif state_type == sr.StateType.JOINT_TRAJECTORY: + return clproto.MessageType.JOINT_TRAJECTORY_MESSAGE return clproto.MessageType.UNKNOWN_MESSAGE @@ -155,10 +160,24 @@ def write_stamped_message(message: MsgT, state: StateT, time: rclpy.time.Time): write_message(message.twist, state) elif isinstance(message, geometry.WrenchStamped): write_message(message.wrench, state) + elif isinstance(message, trajectory.JointTrajectory) and isinstance(state, sr.JointTrajectory): + message.joint_names = state.get_joint_names() + point = trajectory.JointTrajectoryPoint() + for i, point in enumerate(state.get_points()): + ros_point = trajectory.JointTrajectoryPoint() + ros_point.positions = point.get_positions().tolist() + ros_point.velocities = point.get_velocities().tolist() + ros_point.accelerations = point.get_accelerations().tolist() + ros_point.effort = point.get_torques().tolist() + ros_point.time_from_start = rclpy.time.Duration( + seconds=state.get_time_from_start(i).total_seconds()).to_msg() + message.points.append(ros_point) + message.header.frame_id = state.get_name() else: raise MessageTranslationError("The provided combination of state type and message type is not supported") message.header.stamp = time.to_msg() - message.header.frame_id = state.get_reference_frame() + if not isinstance(message, trajectory.JointTrajectory): + message.header.frame_id = state.get_reference_frame() except MessageTranslationError: raise except Exception as e: diff --git a/source/modulo_core/package.xml b/source/modulo_core/package.xml index 9ea4af095..a6326b1b3 100644 --- a/source/modulo_core/package.xml +++ b/source/modulo_core/package.xml @@ -20,6 +20,7 @@ std_msgs tf2_msgs modulo_interfaces + trajectory_msgs ament_lint_auto ament_lint_common diff --git a/source/modulo_core/src/translators/message_readers.cpp b/source/modulo_core/src/translators/message_readers.cpp index 3bd9bcf78..a360e9854 100644 --- a/source/modulo_core/src/translators/message_readers.cpp +++ b/source/modulo_core/src/translators/message_readers.cpp @@ -1,4 +1,6 @@ #include "modulo_core/translators/message_readers.hpp" +#include +#include namespace modulo_core::translators { @@ -82,6 +84,41 @@ void read_message(state_representation::JointState& state, const sensor_msgs::ms } } +void read_message(state_representation::JointTrajectory& state, const trajectory_msgs::msg::JointTrajectory& message) { + try { + state.set_name(message.header.frame_id); + if (!message.joint_names.empty()) { + state.set_joint_names(message.joint_names); + } else { + throw std::runtime_error("JointTrajectory message has no joint names"); + } + std::chrono::nanoseconds time_from_start(0); + for (unsigned int i = 0; i < message.points.size(); ++i) { + state_representation::JointState point("point_" + std::to_string(i), state.get_joint_names()); + if (!message.points[i].positions.empty()) { + point.set_positions(message.points[i].positions); + } + if (!message.points[i].velocities.empty()) { + point.set_velocities(message.points[i].velocities); + } + if (!message.points[i].accelerations.empty()) { + point.set_accelerations(message.points[i].accelerations); + } + if (!message.points[i].effort.empty()) { + point.set_torques(message.points[i].effort); + } + auto ros_time_from_start = std::chrono::nanoseconds(message.points[i].time_from_start.nanosec); + auto duration = ros_time_from_start - time_from_start; + time_from_start = ros_time_from_start; + state.add_point(point, duration); + } + } catch (const std::exception& ex) { + throw exceptions::MessageTranslationException(ex.what()); + } catch (...) { + throw exceptions::MessageTranslationException("Unknown error while reading JointTrajectory message"); + } +} + void read_message(bool& state, const std_msgs::msg::Bool& message) { state = message.data; } diff --git a/source/modulo_core/src/translators/message_writers.cpp b/source/modulo_core/src/translators/message_writers.cpp index c7014c21c..3127fe649 100644 --- a/source/modulo_core/src/translators/message_writers.cpp +++ b/source/modulo_core/src/translators/message_writers.cpp @@ -1,6 +1,8 @@ #include "modulo_core/translators/message_writers.hpp" +#include "trajectory_msgs/msg/joint_trajectory.hpp" #include +#include using namespace state_representation; @@ -127,6 +129,33 @@ void write_message(tf2_msgs::msg::TFMessage& message, const CartesianState& stat message.transforms.push_back(transform); } +void write_message( + trajectory_msgs::msg::JointTrajectory& message, const JointTrajectory& state, const rclcpp::Time& time) { + if (!state) { + throw exceptions::MessageTranslationException( + state.get_name() + " state is empty while attempting to write it to message"); + } + message.set__joint_names(state.get_joint_names()); + for (unsigned int i = 0; i < state.get_size(); ++i) { + auto [joint_states, duration] = state[i]; + trajectory_msgs::msg::JointTrajectoryPoint point; + point.positions.assign( + joint_states.get_positions().data(), joint_states.get_positions().data() + joint_states.get_positions().size()); + point.velocities.assign( + joint_states.get_velocities().data(), + joint_states.get_velocities().data() + joint_states.get_velocities().size()); + point.accelerations.assign( + joint_states.get_accelerations().data(), + joint_states.get_accelerations().data() + joint_states.get_accelerations().size()); + point.effort.assign( + joint_states.get_torques().data(), joint_states.get_torques().data() + joint_states.get_torques().size()); + point.time_from_start = rclcpp::Duration(state.get_time_from_start(i)); + message.points.push_back(point); + } + message.header.stamp = time; + message.header.frame_id = state.get_name(); +} + template void write_message(U& message, const Parameter& state, const rclcpp::Time&) { if (!state) { diff --git a/source/modulo_core/test/cpp/translators/test_messages.cpp b/source/modulo_core/test/cpp/translators/test_messages.cpp index 765bf0b40..9b30c0a9c 100644 --- a/source/modulo_core/test/cpp/translators/test_messages.cpp +++ b/source/modulo_core/test/cpp/translators/test_messages.cpp @@ -1,9 +1,13 @@ +#include #include #include "modulo_core/translators/message_readers.hpp" #include "modulo_core/translators/message_writers.hpp" +#include "trajectory_msgs/msg/joint_trajectory.hpp" #include +#include +#include #include using namespace modulo_core::translators; @@ -42,9 +46,20 @@ class MessageTranslatorsTest : public ::testing::Test { void SetUp() override { state_ = state_representation::CartesianState::Random("test", "reference"); joint_state_ = state_representation::JointState::Random("robot", 3); + ctrajectory_ = state_representation::CartesianTrajectory("test", "reference"); + jtrajectory_.set_joint_names({"joint_0", "joint_1", "joint_2"}); + for (unsigned int i = 0; i < 10; i++) { + ctrajectory_.add_point( + state_representation::CartesianState::Random("test", "reference"), std::chrono::nanoseconds((i + 1) * 10)); + jtrajectory_.add_point( + state_representation::JointState::Random("robot", jtrajectory_.get_joint_names()), + std::chrono::nanoseconds((i + 1) * 10)); + } } state_representation::CartesianState state_; state_representation::JointState joint_state_; + state_representation::CartesianTrajectory ctrajectory_; + state_representation::JointTrajectory jtrajectory_; rclcpp::Clock clock_; }; @@ -256,3 +271,54 @@ TEST_F(MessageTranslatorsTest, TestEncodedStatePointerIncompatibleType) { auto new_state_ptr = state_representation::make_shared_state(state_representation::CartesianPose()); EXPECT_THROW(read_message(new_state_ptr, message), modulo_core::exceptions::MessageTranslationException); } + +TEST_F(MessageTranslatorsTest, TestTrajectory) { + {// CartesianTrajectory + auto state_ptr = state_representation::make_shared_state(ctrajectory_); + auto message = modulo_core::EncodedState(); + write_message(message, state_ptr, clock_.now()); + auto new_state_ptr = state_representation::make_shared_state(state_representation::CartesianTrajectory()); + read_message(new_state_ptr, message); + EXPECT_EQ(new_state_ptr->get_type(), state_representation::StateType::CARTESIAN_TRAJECTORY); + auto new_state = *std::dynamic_pointer_cast(new_state_ptr); + EXPECT_EQ(ctrajectory_.get_size(), new_state.get_size()); + EXPECT_EQ(ctrajectory_.get_name(), new_state.get_name()); + EXPECT_EQ(ctrajectory_.get_reference_frame(), new_state.get_reference_frame()); + for (unsigned int i = 0; i < ctrajectory_.get_size(); ++i) { + EXPECT_TRUE(ctrajectory_.get_point(i).data().isApprox(new_state.get_point(i).data())); + EXPECT_EQ(ctrajectory_.get_duration(i), new_state.get_duration(i)); + } + } + + {// JointTrajectory + auto state_ptr = state_representation::make_shared_state(jtrajectory_); + auto message = modulo_core::EncodedState(); + write_message(message, state_ptr, clock_.now()); + auto new_state_ptr = state_representation::make_shared_state(state_representation::JointTrajectory()); + read_message(new_state_ptr, message); + EXPECT_EQ(new_state_ptr->get_type(), state_representation::StateType::JOINT_TRAJECTORY); + auto new_state = *std::dynamic_pointer_cast(new_state_ptr); + EXPECT_EQ(jtrajectory_.get_size(), new_state.get_size()); + EXPECT_EQ(jtrajectory_.get_name(), new_state.get_name()); + EXPECT_EQ(jtrajectory_.get_joint_names(), new_state.get_joint_names()); + for (unsigned int i = 0; i < jtrajectory_.get_size(); ++i) { + EXPECT_TRUE(jtrajectory_.get_point(i).data().isApprox(new_state.get_point(i).data())); + EXPECT_EQ(jtrajectory_.get_duration(i), new_state.get_duration(i)); + } + } + + {// test trajectory_msgs::msg::JointTrajectory conversion to/from JointTrajectory + trajectory_msgs::msg::JointTrajectory ros_trajectory; + write_message(ros_trajectory, jtrajectory_, clock_.now()); + state_representation::JointTrajectory new_state; + read_message(new_state, ros_trajectory); + + EXPECT_EQ(jtrajectory_.get_size(), new_state.get_size()); + EXPECT_EQ(jtrajectory_.get_name(), new_state.get_name()); + EXPECT_EQ(jtrajectory_.get_joint_names(), new_state.get_joint_names()); + for (unsigned int i = 0; i < jtrajectory_.get_size(); ++i) { + EXPECT_TRUE(jtrajectory_.get_point(i).data().isApprox(new_state.get_point(i).data())); + EXPECT_EQ(jtrajectory_.get_duration(i), new_state.get_duration(i)); + } + } +} diff --git a/source/modulo_core/test/python/conftest.py b/source/modulo_core/test/python/conftest.py index 9f154ee92..2e27160a3 100644 --- a/source/modulo_core/test/python/conftest.py +++ b/source/modulo_core/test/python/conftest.py @@ -2,6 +2,7 @@ import pytest import rclpy.clock import state_representation as sr +import datetime from rclpy import Parameter @@ -23,6 +24,25 @@ def clock(): return rclpy.clock.Clock() +@pytest.fixture +def cartesian_trajectory(): + trajectory = sr.CartesianTrajectory("test", "ref") + for i in range(10): + trajectory.add_point(sr.CartesianState().Random("test", "ref"), datetime.timedelta(seconds=(i+1)*10)) + return trajectory + + +@pytest.fixture +def joint_trajectory(): + trajectory = sr.JointTrajectory("test") + trajectory.set_joint_names(["joint1", "joint2", "joint3"]) + for i in range(10): + trajectory.add_point( + sr.JointState().Random("test", trajectory.get_joint_names()), + datetime.timedelta(seconds=(i + 1.05) * 10)) + return trajectory + + @pytest.fixture def parameters(): return {"bool": [True, sr.ParameterType.BOOL, False, Parameter.Type.BOOL], diff --git a/source/modulo_core/test/python/translators/test_messages.py b/source/modulo_core/test/python/translators/test_messages.py index c1d12f541..eeee73045 100644 --- a/source/modulo_core/test/python/translators/test_messages.py +++ b/source/modulo_core/test/python/translators/test_messages.py @@ -9,6 +9,7 @@ from modulo_core.exceptions import MessageTranslationError from rclpy.clock import Clock from sensor_msgs.msg import JointState +import trajectory_msgs.msg as trajectory def read_xyz(message): @@ -171,3 +172,56 @@ def test_encoded_state(cart_state: sr.CartesianState): assert_np_array_equal(new_state.data(), cart_state.data()) assert new_state.get_name() == cart_state.get_name() assert new_state.get_reference_frame() == cart_state.get_reference_frame() + + +def test_cartesian_trajectory(cartesian_trajectory: sr.CartesianTrajectory): + message = EncodedState() + modulo_writers.write_clproto_message(message, cartesian_trajectory, + clproto.MessageType.CARTESIAN_TRAJECTORY_MESSAGE) + new_state = modulo_readers.read_clproto_message(message) + + assert cartesian_trajectory.get_size() == new_state.get_size() + assert cartesian_trajectory.get_name() == new_state.get_name() + assert cartesian_trajectory.get_reference_frame() == new_state.get_reference_frame() + for [point1, duration1, point2, duration2] in zip( + cartesian_trajectory.get_points(), + cartesian_trajectory.get_durations(), + new_state.get_points(), + new_state.get_durations()): + assert_np_array_equal(point1.data(), point2.data()) + assert duration1 == duration2 + + +def test_joint_trajectory(joint_trajectory: sr.JointTrajectory, clock: Clock): + # test encoded state + message = EncodedState() + modulo_writers.write_clproto_message(message, joint_trajectory, + clproto.MessageType.JOINT_TRAJECTORY_MESSAGE) + new_state = modulo_readers.read_clproto_message(message) + + assert joint_trajectory.get_size() == new_state.get_size() + assert joint_trajectory.get_name() == new_state.get_name() + assert joint_trajectory.get_joint_names() == new_state.get_joint_names() + for [point1, duration1, point2, duration2] in zip( + joint_trajectory.get_points(), + joint_trajectory.get_durations(), + new_state.get_points(), + new_state.get_durations()): + assert_np_array_equal(point1.data(), point2.data()) + assert duration1 == duration2 + + # test trajectory message from ROS trajectory message + message = trajectory.JointTrajectory() + modulo_writers.write_stamped_message(message, joint_trajectory, clock.now()) + new_state = modulo_readers.read_message(sr.JointTrajectory(), message) + + assert joint_trajectory.get_size() == new_state.get_size() + assert joint_trajectory.get_name() == new_state.get_name() + assert joint_trajectory.get_joint_names() == new_state.get_joint_names() + for [point1, duration1, point2, duration2] in zip( + joint_trajectory.get_points(), + joint_trajectory.get_durations(), + new_state.get_points(), + new_state.get_durations()): + assert_np_array_equal(point1.data(), point2.data()) + assert duration1 == duration2