From 1fead65913ac66deb985597560a6a6fcf1f538a0 Mon Sep 17 00:00:00 2001 From: ZGZ713912 Date: Thu, 24 Sep 2026 20:41:09 +0800 Subject: [PATCH 1/3] feat(rl_bridge): implement observation frame stacking --- config/executor.yaml | 9 ++--- src/rl_bridge.cpp | 32 ++++++++++++++---- tool/check_policy_contract.py | 5 +-- tool/rl_layout.py | 62 ++++++++++++++++++++++++++--------- tool/stamp_layout_metadata.py | 3 +- 5 files changed, 80 insertions(+), 31 deletions(-) diff --git a/config/executor.yaml b/config/executor.yaml index 02b0244..98e4fee 100644 --- a/config/executor.yaml +++ b/config/executor.yaml @@ -11,10 +11,6 @@ rmcs_executor: wheel_leg_infantry_rl: ros__parameters: board_serial: "" - hip_kp: 200.0 - hip_kd: 4.0 - knee_kp: 200.0 - knee_kd: 4.0 wheel_leg_chassis_controller: ros__parameters: @@ -29,9 +25,10 @@ wheel_leg_chassis_controller: rl_bridge: ros__parameters: + history_length: 1 rl_base: "/wheel_leg/rl" - policy_rate: 50.0 - max_action_age: 0.04 # 缺省 = 2 / policy_rate;超时 → valid=0(绝不保持旧动作) + policy_rate: 100.0 + max_action_age: 0.02 # 缺省 = 2 / policy_rate;超时 → valid=0(绝不保持旧动作) rl_obs_size: 28 rl_action_size: 6 diff --git a/src/rl_bridge.cpp b/src/rl_bridge.cpp index 76841a2..aac83b4 100644 --- a/src/rl_bridge.cpp +++ b/src/rl_bridge.cpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include #include @@ -284,12 +285,18 @@ class RlBridge final term.index = cursor; cursor += term.dim; } - obs_size_ = cursor; + obs_frame_size_ = cursor; + const auto history_length = integer_or_("history_length").value_or(1); + if (history_length < 1 || history_length > 64) + throw std::invalid_argument("RlBridge: history_length must be in [1, 64]"); + history_length_ = static_cast(history_length); + obs_size_ = obs_frame_size_ * history_length_; if (const auto declared = integer_or_("rl_obs_size"); declared.has_value() && static_cast(*declared) != obs_size_) throw std::invalid_argument( "RlBridge: rl_obs_size=" + std::to_string(*declared) - + " but observation_terms sum to " + std::to_string(obs_size_)); + + " but observation frame/history contract is " + std::to_string(obs_frame_size_) + + "x" + std::to_string(history_length_) + "=" + std::to_string(obs_size_)); policy_rate_ = number_or_("policy_rate", 50.0); if (!(policy_rate_ > 0.0) || !is_finite(policy_rate_)) @@ -420,16 +427,23 @@ class RlBridge final if (read_unsigned_(reset_slot_, reset_count) && reset_count != last_reset_count_) { last_reset_count_ = reset_count; reset_runtime_(); - RCLCPP_INFO( - get_logger(), "RMCS reset: last_action cleared (reset_count=%llu)", - static_cast(reset_count)); } } if (!pub_started_ || now - last_pub_time_ >= pub_period_) { std::vector obs; if (build_observation_(obs)) { - publish_observation_(obs, now); + if (history_.empty()) + history_.assign(history_length_, obs); + else { + history_.pop_front(); + history_.push_back(obs); + } + std::vector stacked; + stacked.reserve(obs_size_); + for (const auto& frame : history_) + stacked.insert(stacked.end(), frame.begin(), frame.end()); + publish_observation_(stacked, now); ++pub_ok_count_; } else { ++obs_invalid_count_; @@ -1162,7 +1176,7 @@ class RlBridge final } std::string obs_layout_signature_() const { - std::string signature = "v2"; + std::string signature = "v3-history=" + std::to_string(history_length_); for (const auto& term : obs_terms_) { signature += "|" + term.id; if (term.scale != 1.0) @@ -1467,6 +1481,7 @@ class RlBridge final void reset_runtime_() { std::fill(last_actions_.begin(), last_actions_.end(), 0.0); std::fill(written_.begin(), written_.end(), 0.0); + history_.clear(); pub_started_ = false; prev_pub_seq_ = 0; } @@ -1503,6 +1518,8 @@ class RlBridge final std::unordered_set own_output_paths_; std::size_t obs_size_ = 0; + std::size_t obs_frame_size_ = 0; + std::size_t history_length_ = 1; std::size_t action_size_ = 0; std::string rl_base_; @@ -1522,6 +1539,7 @@ class RlBridge final std::vector last_actions_; std::vector written_; + std::deque> history_; std::chrono::nanoseconds pub_period_{std::chrono::milliseconds(20)}; std::chrono::steady_clock::time_point last_pub_time_{}; diff --git a/tool/check_policy_contract.py b/tool/check_policy_contract.py index c862d64..3bc113a 100644 --- a/tool/check_policy_contract.py +++ b/tool/check_policy_contract.py @@ -80,7 +80,7 @@ def info(self, name, detail=""): def _print_layout(config, node): obs_terms, act_terms, obs_size, act_size = layout.load_config(config, node) - obs_sig = layout.obs_signature(obs_terms, act_size) + obs_sig = layout.obs_signature(obs_terms, act_size, layout.history_length(config, node)) act_sig = layout.action_signature(act_terms) digest = layout.layout_hash(obs_sig, act_sig, obs_size, act_size) print(f"layout_hash : {layout.hex16(digest)}") @@ -197,7 +197,8 @@ def main() -> None: if args.config: try: obs_terms, act_terms, obs_size, act_size = layout.load_config(args.config, args.node) - obs_sig = layout.obs_signature(obs_terms, act_size) + obs_sig = layout.obs_signature( + obs_terms, act_size, layout.history_length(config, node)) act_sig = layout.action_signature(act_terms) except layout.LayoutError as exc: print(f"FAIL config: {exc}", file=sys.stderr) diff --git a/tool/rl_layout.py b/tool/rl_layout.py index 3cd93ff..acc850e 100644 --- a/tool/rl_layout.py +++ b/tool/rl_layout.py @@ -47,12 +47,14 @@ FNV1A64_PRIME = 0x100000001B3 MASK64 = 0xFFFFFFFFFFFFFFFF -SIGNATURE_VERSION = "v2" +OBS_SIGNATURE_VERSION = "v3" +ACTION_SIGNATURE_VERSION = "v2" DEFAULT_NODE = "rl_bridge" OBS_KEYS = frozenset({ "path", "take", "transform", "type", "index", "scale", "clip", "default", "name", "joints", "relative", "zero", "indices", "value", + "history_length", }) ACTION_KEYS = frozenset({"index", "output", "name", "scale", "clip"}) @@ -423,14 +425,25 @@ def _obs_entry(term: Dict) -> str: return entry + "@" + str(term["dim"]) -def obs_signature(terms: Sequence[str], action_size: int) -> str: +def obs_signature(terms: Sequence[str], action_size: int, history_length: int = 1) -> str: """canonical v2 观测布局串:"v2" ( "|" [*scale]@dim )*。""" + if int(history_length) < 1: + raise LayoutError("history_length 必须 >= 1") entries = [_obs_entry(parse_obs_term(spec, action_size)) for spec in terms] - return "|".join([SIGNATURE_VERSION] + entries) + return "|".join([f"{OBS_SIGNATURE_VERSION}-history={int(history_length)}"] + entries) def obs_signature_dim(signature: str) -> int: """规范串里所有 entry 的 dim 之和(用于「metadata 自洽」检查)。""" + history = 1 + if signature.startswith("v3-history="): + prefix, _, _ = signature.partition("|") + try: + history = int(prefix.split("=", 1)[1]) + except (IndexError, ValueError) as exc: + raise LayoutError(f"历史帧签名非法:{signature!r}") from exc + if history < 1: + raise LayoutError(f"历史帧长度非法:{signature!r}") total = 0 for entry in _signature_entries(signature): _, _, dim_text = entry.rpartition("@") @@ -438,7 +451,7 @@ def obs_signature_dim(signature: str) -> int: total += int(dim_text) except ValueError: raise LayoutError(f"规范串 entry 缺少 @dim({entry!r}):{signature!r}") - return total + return total * history @@ -497,16 +510,16 @@ def action_signature(terms: Sequence[str]) -> str: raise LayoutError( f"动作 index 必须是 0..{len(parsed) - 1} 的完整置换(无空洞/无重复),实际 {indices}" ) - return "|".join([SIGNATURE_VERSION] + [term["id"] for term in parsed]) + return "|".join([ACTION_SIGNATURE_VERSION] + [term["id"] for term in parsed]) def _signature_entries(signature: str) -> List[str]: text = (signature or "").strip() - if text in ("", SIGNATURE_VERSION): + if text in ("", OBS_SIGNATURE_VERSION, ACTION_SIGNATURE_VERSION): return [] - if text.startswith(SIGNATURE_VERSION + "|"): - text = text[len(SIGNATURE_VERSION) + 1:] + if text.startswith("v3-history=") or text.startswith("v2|"): + text = text[text.index("|") + 1:] return [entry for entry in text.split("|") if entry != ""] @@ -518,7 +531,8 @@ def parse_signature(signature: str) -> Dict: 「模型带 v1 metadata,需重盖章」而不是一个困惑的 mismatch。 """ text = (signature or "").strip() - version = "v2" if (text == SIGNATURE_VERSION or text.startswith(SIGNATURE_VERSION + "|")) else "v1" + version = "v3" if text.startswith("v3-history=") else ( + "v2" if text.startswith("v2|") or text == "v2" else "v1") return {"version": version, "entries": _signature_entries(text)} @@ -588,16 +602,34 @@ def load_config(config_path, node: str = DEFAULT_NODE): parsed_obs = [parse_obs_term(spec, act_size) for spec in obs_terms] dim_sum = sum(term["dim"] for term in parsed_obs) declared_obs = _declared_size(params, "rl_obs_size", config_path) - obs_size = dim_sum if declared_obs is None else declared_obs - if declared_obs is not None and declared_obs != dim_sum: + history_length = _declared_size(params, "history_length", config_path) or 1 + if history_length < 1: + raise LayoutError(f"history_length={history_length} 必须 >= 1({config_path})") + obs_size = dim_sum * history_length + if declared_obs is not None and declared_obs != obs_size: raise LayoutError( - f"rl_obs_size={declared_obs} 与观测词条维度之和 {dim_sum} 不一致" + f"rl_obs_size={declared_obs} 与单帧维度 {dim_sum} x history_length {history_length}" + f" = {obs_size} 不一致" f"({config_path} 节点 {node})" ) _check_obs_index_order(obs_terms, parsed_obs, config_path, node) return obs_terms, act_terms, obs_size, act_size +def history_length(config_path, node: str = DEFAULT_NODE) -> int: + yaml = _import_yaml() + try: + with open(str(config_path), "r", encoding="utf-8") as handle: + document = yaml.safe_load(handle) + except (OSError, yaml.YAMLError) as exc: + raise LayoutError(f"无法读取配置 {config_path}: {exc}") + params = _select_params(document, node, config_path) + value = _declared_size(params, "history_length", config_path) or 1 + if value < 1: + raise LayoutError(f"history_length={value} 必须 >= 1({config_path})") + return value + + def _select_params(document, node: str, config_path) -> Dict: if not isinstance(document, dict) or not document: raise LayoutError(f"配置 {config_path} 为空或不是 YAML 映射") @@ -768,7 +800,7 @@ def _self_test() -> None: other = parse_obs_term(f"path=/v take=x type={interface}", 6) assert (other["dim"], other["id"]) == (1, "vec3c:/v:x"), other assert obs_signature(["path=/v take=x"], 6) == obs_signature(["path=/v take=x type=direction_vector"], 6) - assert obs_signature(["path=/v take=x"], 6) == "v2|vec3c:/v:x@1" + assert obs_signature(["path=/v take=x"], 6) == "v3-history=1|vec3c:/v:x@1" for axis in ("x", "y", "z"): assert parse_obs_term(f"path=/v take={axis}", 6)["id"] == f"vec3c:/v:{axis}" @@ -817,8 +849,8 @@ def _self_test() -> None: assert parse_obs_term("path=/v clip=-1:1", 6)["clip"] == (-1.0, 1.0) _expect_error(lambda: parse_obs_term("path=/v clip=1:-1", 6), "clip") - assert obs_signature(["path=/v scale=2 clip=10 index=3 default=0.1"], 6) == "v2|/v*2@1" - assert obs_signature(["path=/v"], 6) == "v2|/v@1" + assert obs_signature(["path=/v scale=2 clip=10 index=3 default=0.1"], 6) == "v3-history=1|/v*2@1" + assert obs_signature(["path=/v"], 6) == "v3-history=1|/v@1" action = parse_action_term("index=0 output=/rl/action/lf0") assert (action["index"], action["output"], action["name"], action["id"]) \ diff --git a/tool/stamp_layout_metadata.py b/tool/stamp_layout_metadata.py index c01d9b3..b67346e 100644 --- a/tool/stamp_layout_metadata.py +++ b/tool/stamp_layout_metadata.py @@ -67,7 +67,8 @@ def main() -> None: try: obs_terms, act_terms, obs_size, act_size = layout.load_config(args.config, args.node) - obs_sig = layout.obs_signature(obs_terms, act_size) + obs_sig = layout.obs_signature( + obs_terms, act_size, layout.history_length(args.config, args.node)) act_sig = layout.action_signature(act_terms) digest = layout.layout_hash(obs_sig, act_sig, obs_size, act_size) From 010be6f6290347f507d4eb7d5bafb9a431f80a17 Mon Sep 17 00:00:00 2001 From: ZGZ713912 Date: Thu, 24 Sep 2026 21:32:52 +0800 Subject: [PATCH 2/3] feat: add joint configuration and observation handling --- CMakeLists.txt | 7 + src/policy_server_launcher.cpp | 16 +- src/rl_bridge.cpp | 1454 ++++----------------------- src/rl_bridge/action_channel.cpp | 55 + src/rl_bridge/action_channel.hpp | 69 ++ src/rl_bridge/interface_binding.cpp | 157 +++ src/rl_bridge/interface_binding.hpp | 26 + src/rl_bridge/joint_config.cpp | 81 ++ src/rl_bridge/joint_config.hpp | 27 + src/rl_bridge/observation.cpp | 156 +++ src/rl_bridge/observation.hpp | 18 + src/rl_bridge/parameters.cpp | 99 ++ src/rl_bridge/parameters.hpp | 20 + src/rl_bridge/term_parser.cpp | 285 ++++++ src/rl_bridge/term_parser.hpp | 24 + src/rl_bridge/types.hpp | 104 ++ src/rl_bridge/utility.cpp | 225 +++++ src/rl_bridge/utility.hpp | 44 + tool/check_policy_contract.py | 2 +- 19 files changed, 1604 insertions(+), 1265 deletions(-) create mode 100644 src/rl_bridge/action_channel.cpp create mode 100644 src/rl_bridge/action_channel.hpp create mode 100644 src/rl_bridge/interface_binding.cpp create mode 100644 src/rl_bridge/interface_binding.hpp create mode 100644 src/rl_bridge/joint_config.cpp create mode 100644 src/rl_bridge/joint_config.hpp create mode 100644 src/rl_bridge/observation.cpp create mode 100644 src/rl_bridge/observation.hpp create mode 100644 src/rl_bridge/parameters.cpp create mode 100644 src/rl_bridge/parameters.hpp create mode 100644 src/rl_bridge/term_parser.cpp create mode 100644 src/rl_bridge/term_parser.hpp create mode 100644 src/rl_bridge/types.hpp create mode 100644 src/rl_bridge/utility.cpp create mode 100644 src/rl_bridge/utility.hpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 91bd886..5752884 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -53,6 +53,13 @@ include_directories(${PROJECT_SOURCE_DIR}/src) ament_auto_add_library(rmcs_rl_bridge SHARED src/rl_bridge.cpp + src/rl_bridge/utility.cpp + src/rl_bridge/parameters.cpp + src/rl_bridge/joint_config.cpp + src/rl_bridge/term_parser.cpp + src/rl_bridge/interface_binding.cpp + src/rl_bridge/observation.cpp + src/rl_bridge/action_channel.cpp src/policy_server_launcher.cpp ) target_link_libraries(rmcs_rl_bridge ${cpp_typesupport_target}) diff --git a/src/policy_server_launcher.cpp b/src/policy_server_launcher.cpp index a22561d..baf4677 100644 --- a/src/policy_server_launcher.cpp +++ b/src/policy_server_launcher.cpp @@ -1,4 +1,3 @@ - #include #include #include @@ -21,15 +20,15 @@ namespace rmcs_rl { -// 随 executor 生命周期拉起独立的 policy_server 子进程: -// - 组件只存在于需要 RL 的配置里,非 RL 车不受影响; -// - 子进程设置 PR_SET_PDEATHSIG,executor 结束/崩溃时自动被内核回收,不会留孤儿; -// - update() 低频 waitpid(WNOHANG) 监管,可选退避重启。 +// Spawns and supervises an independent policy_server child process for the executor lifetime: +// - the component only exists in RL configs; non-RL robots are unaffected; +// - the child sets PR_SET_PDEATHSIG so the kernel reaps it when executor exits/crashes; +// - update() does a low-frequency waitpid(WNOHANG) poll with optional backoff restart. class PolicyServerLauncher : public rmcs_executor::Component , public rclcpp::Node { public: - PolicyServerLauncher() + explicit PolicyServerLauncher() : Node( get_component_name(), rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { @@ -168,7 +167,7 @@ class PolicyServerLauncher } if (pid == 0) { - // 子进程内只调用 async-signal-safe 的接口。 + // Only async-signal-safe calls are allowed inside the child. ::prctl(PR_SET_PDEATHSIG, SIGTERM); if (::getppid() != parent_pid) ::_exit(EXIT_FAILURE); @@ -230,8 +229,7 @@ class PolicyServerLauncher ::kill(pid, SIGKILL); ::waitpid(pid, nullptr, 0); - RCLCPP_WARN( - get_logger(), "policy_server (pid=%d) did not stop, killed", static_cast(pid)); + RCLCPP_WARN(get_logger(), "policy_server (pid=%d) did not stop, killed", static_cast(pid)); } bool autostart_ = true; diff --git a/src/rl_bridge.cpp b/src/rl_bridge.cpp index aac83b4..3c13d83 100644 --- a/src/rl_bridge.cpp +++ b/src/rl_bridge.cpp @@ -1,237 +1,153 @@ - #include -#include #include #include #include #include -#include #include #include -#include #include -#include #include #include #include #include -#include -#include #include #include #include -#include #include #include #include #include #include -#include #include #include #include +#include "rl_bridge/action_channel.hpp" +#include "rl_bridge/interface_binding.hpp" +#include "rl_bridge/joint_config.hpp" +#include "rl_bridge/observation.hpp" +#include "rl_bridge/parameters.hpp" +#include "rl_bridge/term_parser.hpp" +#include "rl_bridge/types.hpp" +#include "rl_bridge/utility.hpp" #include "rl_layout.hpp" namespace rmcs_rl { -namespace { - -constexpr std::size_t kNoSlot = std::numeric_limits::max(); - -std::vector split_by(const std::string& text, char delimiter) { - std::vector parts; - std::string current; - std::istringstream stream{text}; - while (std::getline(stream, current, delimiter)) - parts.push_back(current); - return parts; -} - -std::vector split_whitespace(const std::string& text) { - std::vector tokens; - std::istringstream stream{text}; - std::string token; - while (stream >> token) - tokens.push_back(token); - return tokens; -} - -std::string format_number(double value) { - char buffer[32]; - std::snprintf(buffer, sizeof(buffer), "%.6g", value); - return buffer; -} - -std::string join(const std::vector& parts, const char* separator) { - std::string text; - for (std::size_t i = 0; i < parts.size(); ++i) { - if (i != 0) - text += separator; - text += parts[i]; - } - return text; -} - -double parse_double(const std::string& text, const std::string& context) { - try { - std::size_t consumed = 0; - const double value = std::stod(text, &consumed); - if (consumed != text.size()) - throw std::invalid_argument("trailing characters"); - return value; - } catch (const std::exception&) { - throw std::invalid_argument("RlBridge: '" + text + "' is not a number (" + context + ")"); - } -} - -std::int64_t parse_integer(const std::string& text, const std::string& context) { - try { - std::size_t consumed = 0; - const auto value = std::stoll(text, &consumed); - if (consumed != text.size()) - throw std::invalid_argument("trailing characters"); - return value; - } catch (const std::exception&) { - throw std::invalid_argument("RlBridge: '" + text + "' is not an integer (" + context + ")"); - } -} - -bool parse_boolean(const std::string& text, const std::string& context) { - if (text == "true" || text == "1") - return true; - if (text == "false" || text == "0") - return false; - throw std::invalid_argument( - "RlBridge: '" + text + "' is not a boolean (true/false) (" + context + ")"); -} - -bool is_finite(double value) { return std::isfinite(value); } - -std::string pretty_type(const std::type_info& type) { - if (type == typeid(double)) - return "double"; - if (type == typeid(float)) - return "float"; - if (type == typeid(bool)) - return "bool"; - if (type == typeid(int)) - return "int"; - if (type == typeid(std::size_t)) - return "std::size_t"; - if (type == typeid(Eigen::Vector3d)) - return "Eigen::Vector3d"; - if (type == typeid(rmcs_description::BaseLink::DirectionVector)) - return "rmcs_description::BaseLink::DirectionVector"; - if (type == typeid(Eigen::Quaterniond)) - return "Eigen::Quaterniond"; - return std::string{type.name()} + " (unknown)"; -} - -} // namespace - class RlBridge final : public rmcs_executor::Component , public rclcpp::Node { +public: + explicit RlBridge() + : Node( + get_component_name(), + rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { + + rl_base_ = string_or(*this, "rl_base", "/rl"); + + joint_config_ = load_joint_config(*this); - enum class TermKind { kPath, kJointPos, kJointVel, kJointTorque, kLastAction, kConstant }; + load_action_terms_(); + load_observation_terms_(); + load_runtime_parameters_(); - enum class Take { kScalar, kComponent, kVector, kGravity }; + register_status_outputs_(); + setup_topics_(); - enum class Binding { - kDouble, - kBool, - kInt, - kSize, - kVector3, - kDirectionVector, - kQuaternion, - }; + action_channel_.resize(action_size_); + read_snapshot_.action.assign(action_size_, 0.0); + } - struct ObsTerm { - TermKind kind = TermKind::kPath; - Take take = Take::kScalar; + void before_pairing(const OutputInfoMap& output_map) override { + bind_observation_slots_(output_map); + bind_control_slots_(output_map); - std::string path; - std::string id; - std::size_t dim = 1; + for (const auto& slot : slots_) + if (own_output_paths_.count(slot.path) != 0) + throw std::runtime_error( + "RlBridge: interface \"" + slot.path + + "\" is produced by RlBridge itself (self reference)"); - std::size_t index = 0; - bool has_index = false; - int component = 0; - bool has_binding = false; - Binding binding = Binding::kDouble; + obs_signature_ = obs_layout_signature(obs_terms_, history_length_); + actions_signature_ = actions_layout_signature(action_terms_); + layout_hash_ = + rmcs_rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); - std::vector joints; - std::vector joint_names; - std::vector joint_slots; - std::vector joint_defaults; - bool relative = false; - bool zero = false; + log_layout_(); + RCLCPP_INFO( + get_logger(), "contract: obs_size=%zu action_size=%zu policy_rate=%.3f Hz", obs_size_, + action_size_, policy_rate_); + RCLCPP_INFO( + get_logger(), "layout_hash=%s (obs/action signatures below)", + hex16(layout_hash_).c_str()); + RCLCPP_INFO(get_logger(), " obs signature : %s", obs_signature_.c_str()); + RCLCPP_INFO(get_logger(), " action signature : %s", actions_signature_.c_str()); + RCLCPP_INFO( + get_logger(), "topics: %s/obs -> %s/action (best_effort, keep_last=1)", + rl_base_.c_str(), rl_base_.c_str()); + } - std::vector action_indices; - std::vector constants; + void update() override { + const auto now = std::chrono::steady_clock::now(); - double scale = 1.0; - bool has_clip = false; - double clip_min = 0.0; - double clip_max = 0.0; - bool has_default = false; - double default_value = 0.0; + maybe_handle_reset_(); + maybe_publish_observation_(now); - std::size_t slot = kNoSlot; - }; + ActionSnapshot& snapshot = read_snapshot_; + const bool has_snapshot = action_channel_.try_read(snapshot); - struct ActionTerm { - std::size_t index = 0; - std::string output; - std::string id; - double scale = 1.0; - bool has_clip = false; - double clip_min = 0.0; - double clip_max = 0.0; - }; + double age = std::numeric_limits::quiet_NaN(); + bool seq_ok = false; + bool finite = true; + if (has_snapshot) { + age = obs_age_of(snapshot.obs_seq, now); + seq_ok = (snapshot.obs_seq == pub_seq_) || (snapshot.obs_seq + 1 == pub_seq_); + finite = std::all_of(snapshot.action.begin(), snapshot.action.end(), [](double value) { + return std::isfinite(value); + }); + check_contract_(snapshot); + } - struct Slot { - std::string path; - Binding binding = Binding::kDouble; - bool required = true; - std::unique_ptr> double_value; - std::unique_ptr> bool_value; - std::unique_ptr> int_value; - std::unique_ptr> size_value; - std::unique_ptr> vector3_value; - std::unique_ptr> - direction_vector_value; - std::unique_ptr> quaternion_value; - }; + const bool enabled = read_enable_(); + const bool fresh = has_snapshot && is_finite(age) && age <= max_action_age_; + const bool valid = enabled && contract_ok_ && has_snapshot && fresh && seq_ok && finite; - struct ActionSnapshot { - std::uint64_t obs_seq = 0; - std::uint64_t layout_hash = 0; - std::uint64_t model_id = 0; - std::vector action; - std::chrono::steady_clock::time_point received{}; - }; + write_actions(valid, snapshot, action_terms_, invalid_mode_, written_, action_outputs_); -public: - explicit RlBridge() - : Node( - get_component_name(), - rclcpp::NodeOptions{}.automatically_declare_parameters_from_overrides(true)) { + if (valid) + last_actions_ = snapshot.action; - rl_base_ = string_or_("rl_base", "/rl"); + *valid_output_ = valid ? 1.0 : 0.0; + *healthy_output_ = (contract_ok_ && fresh) ? 1.0 : 0.0; + *action_age_output_ = age; + *obs_seq_output_ = pub_seq_; - load_joint_config_(); + if (valid != last_valid_) { + if (valid) { + RCLCPP_INFO( + get_logger(), "valid=1 (obs_seq=%llu model_id=%s action_age=%.4f s)", + static_cast(snapshot.obs_seq), + hex16(snapshot.model_id).c_str(), age); + } else { + RCLCPP_WARN( + get_logger(), "valid=0: %s", + invalid_reason(enabled, contract_ok_, has_snapshot, fresh, seq_ok, finite) + .c_str()); + } + last_valid_ = valid; + } + } - const auto action_specs = string_array_or_("action_terms"); +private: + void load_action_terms_() { + const auto action_specs = string_array_or(*this, "action_terms"); if (action_specs.empty()) throw std::invalid_argument("RlBridge: required parameter 'action_terms' is missing"); std::unordered_set output_paths; std::set action_indices; for (const auto& spec : action_specs) { - auto term = parse_action_term_(spec); + auto term = parse_action_term(spec); if (!output_paths.emplace(term.output).second) throw std::invalid_argument( "RlBridge: two action terms map to the same output " @@ -251,7 +167,7 @@ class RlBridge final + std::to_string(action_terms_.size() - 1) + "; got index=" + std::to_string(action_terms_[i].index) + " at position " + std::to_string(i)); action_size_ = action_terms_.size(); - if (const auto declared = integer_or_("rl_action_size"); + if (const auto declared = integer_parameter(*this, "rl_action_size"); declared.has_value() && static_cast(*declared) != action_size_) throw std::invalid_argument( "RlBridge: rl_action_size=" + std::to_string(*declared) + " but action_terms has " @@ -264,51 +180,50 @@ class RlBridge final register_output(term.output, *action_outputs_.back(), 0.0); own_output_paths_.insert(term.output); } + } - const auto observation_specs = string_array_or_("observation_terms"); + void load_observation_terms_() { + const auto observation_specs = string_array_or(*this, "observation_terms"); if (observation_specs.empty()) throw std::invalid_argument( "RlBridge: required parameter 'observation_terms' is " "missing"); + const TermParseContext context{*this, joint_config_, action_size_}; for (const auto& spec : observation_specs) - obs_terms_.push_back(parse_obs_term_(spec)); - - std::size_t cursor = 0; - for (auto& term : obs_terms_) { - if (term.has_index && term.index != cursor) - throw std::invalid_argument( - "RlBridge: observation term '" + term.id - + "' declares " - "index=" - + std::to_string(term.index) + " but the running offset is " - + std::to_string(cursor) + " (index= is an assertion, not a reorder)"); - term.index = cursor; - cursor += term.dim; - } - obs_frame_size_ = cursor; - const auto history_length = integer_or_("history_length").value_or(1); + obs_terms_.push_back(parse_obs_term(spec, context)); + + assign_observation_indices(obs_terms_); + obs_frame_size_ = [&] { + std::size_t cursor = 0; + for (const auto& term : obs_terms_) + cursor += term.dim; + return cursor; + }(); + const auto history_length = integer_parameter(*this, "history_length").value_or(1); if (history_length < 1 || history_length > 64) throw std::invalid_argument("RlBridge: history_length must be in [1, 64]"); history_length_ = static_cast(history_length); obs_size_ = obs_frame_size_ * history_length_; - if (const auto declared = integer_or_("rl_obs_size"); + if (const auto declared = integer_parameter(*this, "rl_obs_size"); declared.has_value() && static_cast(*declared) != obs_size_) throw std::invalid_argument( "RlBridge: rl_obs_size=" + std::to_string(*declared) + " but observation frame/history contract is " + std::to_string(obs_frame_size_) + "x" + std::to_string(history_length_) + "=" + std::to_string(obs_size_)); + } - policy_rate_ = number_or_("policy_rate", 50.0); + void load_runtime_parameters_() { + policy_rate_ = number_or(*this, "policy_rate", 50.0); if (!(policy_rate_ > 0.0) || !is_finite(policy_rate_)) throw std::invalid_argument("RlBridge: policy_rate must be finite and > 0"); pub_period_ = std::chrono::duration_cast( std::chrono::duration(1.0 / policy_rate_)); - max_action_age_ = number_or_("max_action_age", 2.0 / policy_rate_); + max_action_age_ = number_or(*this, "max_action_age", 2.0 / policy_rate_); if (!(max_action_age_ > 0.0) || !is_finite(max_action_age_)) throw std::invalid_argument("RlBridge: max_action_age must be finite and > 0"); - expected_model_id_ = parse_u64_(string_or_("expected_model_id", "0"), "expected_model_id"); + expected_model_id_ = parse_u64(string_or(*this, "expected_model_id", "0"), "expected_model_id"); - const std::string invalid = string_or_("invalid_value", "nan"); + const std::string invalid = string_or(*this, "invalid_value", "nan"); if (invalid == "nan") invalid_mode_ = InvalidMode::kNaN; else if (invalid == "zero") @@ -320,12 +235,14 @@ class RlBridge final "RlBridge: invalid_value must be nan|zero|hold " "(quote it in YAML: invalid_value: \"nan\" — an unquoted " "nan is parsed as a float), got '" - + string_or_("invalid_value", "") + "'"); + + string_or(*this, "invalid_value", "") + "'"); - enable_path_ = string_or_("enable_interface", rl_base_ + "/enable"); - enable_default_ = bool_or_("enable_default", false); - reset_path_ = string_or_("reset_interface", ""); + enable_path_ = string_or(*this, "enable_interface", rl_base_ + "/enable"); + enable_default_ = bool_or(*this, "enable_default", false); + reset_path_ = string_or(*this, "reset_interface", ""); + } + void register_status_outputs_() { register_output(rl_base_ + "/valid", valid_output_, 0.0); register_output(rl_base_ + "/healthy", healthy_output_, 0.0); register_output( @@ -335,18 +252,17 @@ class RlBridge final own_output_paths_.insert(rl_base_ + "/healthy"); own_output_paths_.insert(rl_base_ + "/action_age"); own_output_paths_.insert(rl_base_ + "/obs_seq"); + } + void setup_topics_() { obs_publisher_ = create_publisher( rl_base_ + "/obs", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort()); action_subscription_ = create_subscription( rl_base_ + "/action", rclcpp::QoS{rclcpp::KeepLast(1)}.best_effort(), [this](rmcs_rl::msg::Action::UniquePtr message) { on_action_(std::move(message)); }); - - incoming_.action.assign(action_size_, 0.0); - read_snapshot_.action.assign(action_size_, 0.0); } - void before_pairing(const OutputInfoMap& output_map) override { + void bind_observation_slots_(const OutputInfoMap& output_map) { for (auto& term : obs_terms_) { switch (term.kind) { case TermKind::kConstant: @@ -355,13 +271,13 @@ class RlBridge final case TermKind::kJointVel: case TermKind::kJointTorque: { for (std::size_t j = 0; j < term.joints.size(); ++j) { - const auto& joint = joint_names_[term.joints[j]]; + const auto& joint = joint_config_.names[term.joints[j]]; const char field = term.kind == TermKind::kJointPos ? 'a' : term.kind == TermKind::kJointVel ? 'v' : 't'; - const auto path = joint_path_(joint, field); - term.joint_slots.push_back( - acquire_slot_(path, {Binding::kDouble}, true, output_map, "joint term")); + const auto path = joint_path(joint_config_, joint, field); + term.joint_slots.push_back(acquire_slot( + *this, path, {Binding::kDouble}, true, output_map, slots_, "joint term")); } break; } @@ -377,823 +293,92 @@ class RlBridge final break; case Take::kGravity: candidates = {Binding::kQuaternion}; break; } - term.slot = acquire_slot_( - term.path, candidates, !term.has_default, output_map, "observation term"); + term.slot = acquire_slot( + *this, term.path, candidates, !term.has_default, output_map, slots_, + "observation term"); break; } } } + } + void bind_control_slots_(const OutputInfoMap& output_map) { if (!enable_path_.empty()) - enable_slot_ = acquire_slot_( - enable_path_, {Binding::kBool, Binding::kDouble}, false, output_map, + enable_slot_ = acquire_slot( + *this, enable_path_, {Binding::kBool, Binding::kDouble}, false, output_map, slots_, "enable interface"); if (!reset_path_.empty()) - reset_slot_ = acquire_slot_( - reset_path_, {Binding::kSize, Binding::kInt, Binding::kDouble}, false, output_map, - "reset interface"); - - for (const auto& slot : slots_) - if (own_output_paths_.count(slot->path) != 0) - throw std::runtime_error( - "RlBridge: interface \"" + slot->path - + "\" is produced by RlBridge itself (self reference)"); - - obs_signature_ = obs_layout_signature_(); - actions_signature_ = actions_layout_signature_(); - layout_hash_ = - rmcs_rl::layout_hash(obs_signature_, actions_signature_, obs_size_, action_size_); - - log_layout_(); - RCLCPP_INFO( - get_logger(), "contract: obs_size=%zu action_size=%zu policy_rate=%.3f Hz", obs_size_, - action_size_, policy_rate_); - RCLCPP_INFO( - get_logger(), "layout_hash=%s (obs/action signatures below)", - hex16(layout_hash_).c_str()); - RCLCPP_INFO(get_logger(), " obs signature : %s", obs_signature_.c_str()); - RCLCPP_INFO(get_logger(), " action signature : %s", actions_signature_.c_str()); - RCLCPP_INFO( - get_logger(), "topics: %s/obs -> %s/action (best_effort, keep_last=1)", - rl_base_.c_str(), rl_base_.c_str()); - } - - void update() override { - const auto now = std::chrono::steady_clock::now(); - - if (reset_slot_ != kNoSlot) { - std::uint64_t reset_count = 0; - if (read_unsigned_(reset_slot_, reset_count) && reset_count != last_reset_count_) { - last_reset_count_ = reset_count; - reset_runtime_(); - } - } - - if (!pub_started_ || now - last_pub_time_ >= pub_period_) { - std::vector obs; - if (build_observation_(obs)) { - if (history_.empty()) - history_.assign(history_length_, obs); - else { - history_.pop_front(); - history_.push_back(obs); - } - std::vector stacked; - stacked.reserve(obs_size_); - for (const auto& frame : history_) - stacked.insert(stacked.end(), frame.begin(), frame.end()); - publish_observation_(stacked, now); - ++pub_ok_count_; - } else { - ++obs_invalid_count_; - RCLCPP_WARN_THROTTLE( - get_logger(), *get_clock(), 1000, - "observation sources not ready or non-finite; no obs published (%llu frames " - "skipped)", - static_cast(obs_invalid_count_)); - } - } - - ActionSnapshot& snapshot = read_snapshot_; - const bool has_snapshot = try_read_action_(snapshot); - - double age = std::numeric_limits::quiet_NaN(); - bool seq_ok = false; - bool finite = true; - if (has_snapshot) { - age = obs_age_of_(snapshot.obs_seq, now); - seq_ok = (snapshot.obs_seq == pub_seq_) || (snapshot.obs_seq + 1 == pub_seq_); - finite = std::all_of(snapshot.action.begin(), snapshot.action.end(), [](double value) { - return std::isfinite(value); - }); - if (snapshot.layout_hash != layout_hash_) { - if (contract_ok_) { - contract_ok_ = false; - RCLCPP_FATAL( - get_logger(), - "contract mismatch: action layout_hash=%s != bridge layout_hash=%s " - "(policy process and bridge disagree; refusing to output actions)", - hex16(snapshot.layout_hash).c_str(), hex16(layout_hash_).c_str()); - } - } else if (expected_model_id_ != 0 && snapshot.model_id != expected_model_id_) { - if (contract_ok_) { - contract_ok_ = false; - RCLCPP_FATAL( - get_logger(), "model mismatch: action model_id=%s != expected_model_id=%s", - hex16(snapshot.model_id).c_str(), hex16(expected_model_id_).c_str()); - } - } - } - - const bool enabled = read_enable_(); - const bool fresh = has_snapshot && is_finite(age) && age <= max_action_age_; - const bool valid = enabled && contract_ok_ && has_snapshot && fresh && seq_ok && finite; - - write_actions_(valid, snapshot); - - if (valid) - last_actions_ = snapshot.action; - - *valid_output_ = valid ? 1.0 : 0.0; - *healthy_output_ = (contract_ok_ && fresh) ? 1.0 : 0.0; - *action_age_output_ = age; - *obs_seq_output_ = pub_seq_; - - if (valid != last_valid_) { - if (valid) { - RCLCPP_INFO( - get_logger(), "valid=1 (obs_seq=%llu model_id=%s action_age=%.4f s)", - static_cast(snapshot.obs_seq), - hex16(snapshot.model_id).c_str(), age); - } else { - RCLCPP_WARN( - get_logger(), "valid=0: %s", - invalid_reason_(enabled, contract_ok_, has_snapshot, fresh, seq_ok, finite) - .c_str()); - } - last_valid_ = valid; - } - } - -private: - enum class InvalidMode { kNaN, kZero, kHold }; - - std::optional number_(const std::string& name) { - rclcpp::Parameter parameter; - try { - if (!get_parameter(name, parameter)) - return std::nullopt; - } catch (const std::exception&) { - return std::nullopt; - } - switch (parameter.get_type()) { - case rclcpp::ParameterType::PARAMETER_DOUBLE: return parameter.as_double(); - case rclcpp::ParameterType::PARAMETER_INTEGER: - return static_cast(parameter.as_int()); - case rclcpp::ParameterType::PARAMETER_NOT_SET: return std::nullopt; - default: throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a number"); - } - } - - double number_or_(const std::string& name, double fallback) { - const auto value = number_(name); - if (!value.has_value()) - return fallback; - if (!is_finite(*value)) - throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); - return *value; - } - - static std::uint64_t parse_u64_(const std::string& text, const std::string& name) { - try { - std::size_t consumed = 0; - const bool hex = text.rfind("0x", 0) == 0 || text.rfind("0X", 0) == 0; - const auto value = std::stoull(hex ? text.substr(2) : text, &consumed, hex ? 16 : 10); - if (consumed != (hex ? text.size() - 2 : text.size())) - throw std::invalid_argument("trailing characters"); - return static_cast(value); - } catch (const std::exception&) { - throw std::invalid_argument( - "RlBridge: parameter '" + name - + "' must be a decimal or 0x-prefixed 64-bit id, got '" + text + "'"); - } + reset_slot_ = acquire_slot( + *this, reset_path_, {Binding::kSize, Binding::kInt, Binding::kDouble}, false, + output_map, slots_, "reset interface"); } - std::optional integer_or_(const std::string& name) { - const auto value = number_(name); - if (!value.has_value()) - return std::nullopt; - const auto rounded = std::llround(*value); - if (std::abs(*value - static_cast(rounded)) > 1e-9) - throw std::invalid_argument("RlBridge: parameter '" + name + "' must be an integer"); - return rounded; - } - - bool bool_or_(const std::string& name, bool fallback) { - rclcpp::Parameter parameter; - try { - if (!get_parameter(name, parameter)) - return fallback; - } catch (const std::exception&) { - return fallback; - } - if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) - return fallback; - if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_BOOL) - throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a boolean"); - return parameter.as_bool(); - } - - std::string string_or_(const std::string& name, const std::string& fallback) { - rclcpp::Parameter parameter; - try { - if (!get_parameter(name, parameter)) - return fallback; - } catch (const std::exception&) { - return fallback; - } - if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) - return fallback; - if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_STRING) - throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a string"); - return parameter.as_string(); - } - - std::vector string_array_or_(const std::string& name) { - std::vector value; - try { - if (!get_parameter(name, value)) - return {}; - } catch (const std::exception&) { - throw std::invalid_argument( - "RlBridge: parameter '" + name + "' must be a list of strings"); - } - return value; - } - - void load_joint_config_() { - joint_names_ = string_array_or_("joint_names"); - joint_base_path_ = string_or_("joint_base_path", ""); - if (!joint_names_.empty() && joint_base_path_.empty()) - throw std::invalid_argument( - "RlBridge: joint_base_path is required when joint_names is " - "set"); - for (std::size_t i = 0; i < joint_names_.size(); ++i) { - if (joint_names_[i].empty()) - throw std::invalid_argument("RlBridge: joint_names contains an empty name"); - if (!joint_index_.emplace(joint_names_[i], i).second) - throw std::invalid_argument( - "RlBridge: duplicate joint name '" + joint_names_[i] + "'"); - } - - const std::string default_angle = string_or_("joint_angle_suffix", "/angle"); - const std::string default_velocity = string_or_("joint_velocity_suffix", "/velocity"); - const std::string default_torque = string_or_("joint_torque_suffix", "/torque"); - for (const auto& joint : joint_names_) { - angle_suffix_[joint] = default_angle; - velocity_suffix_[joint] = default_velocity; - torque_suffix_[joint] = default_torque; - } - - constexpr const char* prefix = "joint_interface_overrides."; - for (const auto& name : list_parameters({"joint_interface_overrides"}, 10).names) { - if (name.rfind(prefix, 0) != 0) - continue; - const std::string remainder = name.substr(std::string{prefix}.size()); - const auto separator = remainder.find_last_of('.'); - if (separator == std::string::npos) - throw std::invalid_argument( - "RlBridge: joint_interface_overrides entry '" + name - + "' must be ."); - const std::string joint = remainder.substr(0, separator); - const std::string field = remainder.substr(separator + 1); - if (joint_index_.count(joint) == 0) - throw std::invalid_argument( - "RlBridge: joint_interface_overrides references unknown joint '" + joint + "'"); - const std::string value = string_or_(name, ""); - if (field == "angle") - angle_suffix_[joint] = value; - else if (field == "velocity") - velocity_suffix_[joint] = value; - else if (field == "torque") - torque_suffix_[joint] = value; - else - throw std::invalid_argument( - "RlBridge: joint_interface_overrides field '" + field - + "' is not angle/velocity/torque"); - } - } - - std::string joint_path_(const std::string& joint, char field) const { - const auto suffix = field == 'a' ? angle_suffix_.at(joint) - : field == 'v' ? velocity_suffix_.at(joint) - : torque_suffix_.at(joint); - return joint_base_path_ + "/" + joint + suffix; - } - - std::optional default_joint_pos_(std::size_t joint) { - const auto name = "default_joint_pos." + joint_names_[joint]; - const auto value = number_(name); - if (!value.has_value()) - return std::nullopt; - if (!is_finite(*value)) - throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); - return value; - } - - static bool binding_matches_type_(Binding binding, const std::type_info& type) { - switch (binding) { - case Binding::kDouble: return type == typeid(double); - case Binding::kBool: return type == typeid(bool); - case Binding::kInt: return type == typeid(int); - case Binding::kSize: return type == typeid(std::size_t); - case Binding::kVector3: return type == typeid(Eigen::Vector3d); - case Binding::kDirectionVector: - return type == typeid(rmcs_description::BaseLink::DirectionVector); - case Binding::kQuaternion: return type == typeid(Eigen::Quaterniond); - } - return false; - } - - static const char* binding_name_(Binding binding) { - switch (binding) { - case Binding::kDouble: return "double"; - case Binding::kBool: return "bool"; - case Binding::kInt: return "int"; - case Binding::kSize: return "std::size_t"; - case Binding::kVector3: return "Eigen::Vector3d"; - case Binding::kDirectionVector: return "BaseLink::DirectionVector"; - case Binding::kQuaternion: return "Eigen::Quaterniond"; - } - return "unknown"; - } - - std::size_t acquire_slot_( - const std::string& path, const std::vector& candidates, bool required, - const OutputInfoMap& output_map, const char* context) { - const auto output = output_map.find(path); - if (output == output_map.end()) { - if (!required) - return kNoSlot; - throw std::runtime_error( - "RlBridge: required input interface \"" + path - + "\" was not produced by any component (" + context - + "); check the interface path or add default= to make it optional"); - } - if (output->second.kind != rmcs_executor::InterfaceKind::Normal) - throw std::runtime_error( - "RlBridge: input interface \"" + path - + "\" exists but is an Event interface; Normal required (" + context + ")"); - const std::type_info& producer_type = output->second.type.get(); - - for (std::size_t i = 0; i < slots_.size(); ++i) - if (slots_[i]->path == path && binding_matches_type_(slots_[i]->binding, producer_type)) - return i; - - Binding selected = Binding::kDouble; - bool found = false; - for (const auto binding : candidates) - if (binding_matches_type_(binding, producer_type)) { - selected = binding; - found = true; - break; - } - if (!found) { - std::vector expected; - for (const auto binding : candidates) - expected.push_back(binding_name_(binding)); - throw std::runtime_error( - "RlBridge: cannot bind observation interface \"" + path + "\": producer declares \"" - + pretty_type(producer_type) + "\" but the term accepts { " + join(expected, ", ") - + " }. Either fix the term (take=/transform=) or pin type= explicitly."); - } - - auto slot = std::make_unique(); - slot->path = path; - slot->binding = selected; - slot->required = required; - switch (selected) { - case Binding::kDouble: - slot->double_value = std::make_unique>(); - register_input(path, *slot->double_value, required); - break; - case Binding::kBool: - slot->bool_value = std::make_unique>(); - register_input(path, *slot->bool_value, required); - break; - case Binding::kInt: - slot->int_value = std::make_unique>(); - register_input(path, *slot->int_value, required); - break; - case Binding::kSize: - slot->size_value = std::make_unique>(); - register_input(path, *slot->size_value, required); - break; - case Binding::kVector3: - slot->vector3_value = std::make_unique>(); - register_input(path, *slot->vector3_value, required); - break; - case Binding::kDirectionVector: - slot->direction_vector_value = - std::make_unique>(); - register_input(path, *slot->direction_vector_value, required); - break; - case Binding::kQuaternion: - slot->quaternion_value = std::make_unique>(); - register_input(path, *slot->quaternion_value, required); - break; - } - slots_.push_back(std::move(slot)); - return slots_.size() - 1; - } - - static std::map parse_tokens_(const std::string& spec) { - std::map tokens; - for (const auto& token : split_whitespace(spec)) { - const auto equals = token.find('='); - if (equals == std::string::npos || equals == 0) - throw std::invalid_argument( - "RlBridge: term token '" + token + "' must be key=value (term '" + spec + "')"); - const std::string key = token.substr(0, equals); - const std::string value = token.substr(equals + 1); - if (!tokens.emplace(key, value).second) - throw std::invalid_argument( - "RlBridge: duplicate key '" + key + "' in term '" + spec + "'"); - } - return tokens; - } - - static void validate_tokens_( - const std::map& tokens, const std::vector& allowed, - const std::string& spec) { - for (const auto& [key, ignored] : tokens) { - (void)ignored; - if (std::find(allowed.begin(), allowed.end(), key) == allowed.end()) - throw std::invalid_argument( - "RlBridge: unknown key '" + key + "' in term '" + spec + "'"); - } - } - - static double double_token_( - const std::map& tokens, const std::string& key, double fallback, - const std::string& spec) { - const auto token = tokens.find(key); - if (token == tokens.end()) - return fallback; - const double value = parse_double(token->second, key + " in term '" + spec + "'"); - if (!is_finite(value)) - throw std::invalid_argument( - "RlBridge: " + key + " must be finite (term '" + spec + "')"); - return value; - } - - static bool bool_token_( - const std::map& tokens, const std::string& key, bool fallback, - const std::string& spec) { - const auto token = tokens.find(key); - if (token == tokens.end()) - return fallback; - return parse_boolean(token->second, key + " in term '" + spec + "'"); - } - - static std::string string_token_( - const std::map& tokens, const std::string& key, - const std::string& fallback) { - const auto token = tokens.find(key); - if (token == tokens.end()) - return fallback; - return token->second; - } - - static void parse_clip_( - const std::map& tokens, const std::string& spec, bool& has_clip, - double& clip_min, double& clip_max) { - const auto clip = tokens.find("clip"); - if (clip == tokens.end()) + void maybe_handle_reset_() { + if (reset_slot_ == kNoSlot) return; - const auto separator = clip->second.find(':'); - if (separator == std::string::npos) { - const double symmetric = parse_double(clip->second, "clip"); - if (!(symmetric > 0.0) || !is_finite(symmetric)) - throw std::invalid_argument( - "RlBridge: single-value clip must be finite and > 0 (term '" + spec + "')"); - clip_min = -symmetric; - clip_max = symmetric; - } else { - clip_min = parse_double(clip->second.substr(0, separator), "clip min"); - clip_max = parse_double(clip->second.substr(separator + 1), "clip max"); + std::uint64_t reset_count = 0; + if (read_unsigned(slots_[reset_slot_], reset_count) && reset_count != last_reset_count_) { + last_reset_count_ = reset_count; + reset_runtime_(); } - if (!is_finite(clip_min) || !is_finite(clip_max) || clip_min > clip_max) - throw std::invalid_argument("RlBridge: invalid clip range (term '" + spec + "')"); - has_clip = true; } - static void parse_index_( - const std::map& tokens, const std::string& spec, bool& has_index, - std::size_t& index) { - const auto token = tokens.find("index"); - if (token == tokens.end()) + void maybe_publish_observation_(std::chrono::steady_clock::time_point now) { + if (pub_started_ && now - last_pub_time_ < pub_period_) return; - const auto value = parse_integer(token->second, "index in term '" + spec + "'"); - if (value < 0) - throw std::invalid_argument("RlBridge: index must be >= 0 (term '" + spec + "')"); - has_index = true; - index = static_cast(value); - } - static std::vector - parse_index_list_(const std::string& text, std::size_t limit, const std::string& what) { - std::vector indices; - for (const auto& piece : split_by(text, ',')) { - if (piece.empty()) - throw std::invalid_argument("RlBridge: empty entry in " + what + " '" + text + "'"); - const auto range = piece.find(".."); - if (range == std::string::npos) { - const auto index = parse_integer(piece, what); - if (index < 0 || static_cast(index) >= limit) - throw std::invalid_argument( - "RlBridge: " + what + " index " + std::to_string(index) - + " out of range [0, " + std::to_string(limit) + ")"); - indices.push_back(static_cast(index)); - } else { - const auto first = parse_integer(piece.substr(0, range), what); - const auto last = parse_integer(piece.substr(range + 2), what); - if (first < 0 || last < first || static_cast(last) >= limit) - throw std::invalid_argument( - "RlBridge: invalid " + what + " range '" + piece + "' (limit " - + std::to_string(limit) + ")"); - for (auto index = first; index <= last; ++index) - indices.push_back(static_cast(index)); + std::vector obs; + if (build_observation(obs_terms_, slots_, last_actions_, obs_size_, obs)) { + if (history_.empty()) + history_.assign(history_length_, obs); + else { + history_.pop_front(); + history_.push_back(obs); } - } - return indices; - } - - ObsTerm parse_obs_term_(const std::string& spec) { - if (spec.empty()) - throw std::invalid_argument("RlBridge: observation term must not be empty"); - const auto tokens = parse_tokens_(spec); - ObsTerm term; - - std::string type = string_token_(tokens, "type", ""); - - if (type == "joint_pos" || type == "joint_vel" || type == "joint_torque") { - term.kind = type == "joint_pos" ? TermKind::kJointPos - : type == "joint_vel" ? TermKind::kJointVel - : TermKind::kJointTorque; - const auto joints = tokens.find("joints"); - if (joints == tokens.end() || joints->second.empty()) - throw std::invalid_argument( - "RlBridge: joint_* term requires joints=a,b (term '" + spec + "')"); - for (const auto& name : split_by(joints->second, ',')) { - if (name.empty()) - throw std::invalid_argument( - "RlBridge: empty joint name in joints= (term '" + spec + "')"); - const auto iter = joint_index_.find(name); - if (iter == joint_index_.end()) - throw std::invalid_argument( - "RlBridge: unknown joint '" + name + "' (term '" + spec + "')"); - term.joints.push_back(iter->second); - term.joint_names.push_back(name); - } - term.relative = bool_token_(tokens, "relative", false, spec); - term.zero = bool_token_(tokens, "zero", false, spec); - if (term.kind != TermKind::kJointPos && (term.relative || term.zero)) - throw std::invalid_argument( - "RlBridge: relative/zero are only valid for joint_pos (term '" + spec + "')"); - if (term.relative && term.zero) - throw std::invalid_argument( - "RlBridge: joint_pos cannot be both relative and zero (term '" + spec + "')"); - for (const auto joint : term.joints) { - if (term.kind == TermKind::kJointPos && term.relative) { - const auto base = default_joint_pos_(joint); - if (!base.has_value()) - throw std::invalid_argument( - "RlBridge: joint_pos relative term '" + spec - + "' requires default_joint_pos." + joint_names_[joint]); - term.joint_defaults.push_back(*base); - } else { - term.joint_defaults.push_back(0.0); - } - } - term.dim = term.joints.size(); - const std::string flag = term.zero ? "zero" : (term.relative ? "rel" : "abs"); - const std::string prefix = type == "joint_pos" ? "joint_pos" - : type == "joint_vel" ? "joint_vel" - : "joint_torque"; - term.id = prefix + ":" + flag + ":" + join(term.joint_names, "+"); - validate_tokens_( - tokens, {"type", "joints", "relative", "zero", "index", "scale", "clip", "name"}, - spec); - return finish_common_(term, tokens, spec); - } - if (type == "last_action") { - term.kind = TermKind::kLastAction; - if (const auto indices = tokens.find("indices"); indices != tokens.end()) - term.action_indices = - parse_index_list_(indices->second, action_size_, "last_action"); - else - for (std::size_t i = 0; i < action_size_; ++i) - term.action_indices.push_back(i); - term.dim = term.action_indices.size(); - std::vector names; - for (const auto index : term.action_indices) - names.push_back(std::to_string(index)); - term.id = "last_action:" + join(names, "+"); - validate_tokens_(tokens, {"type", "indices", "index", "scale", "clip", "name"}, spec); - return finish_common_(term, tokens, spec); - } - if (type == "constant") { - term.kind = TermKind::kConstant; - const auto value = tokens.find("value"); - if (value == tokens.end() || value->second.empty()) - throw std::invalid_argument( - "RlBridge: constant term requires value=... (term '" + spec + "')"); - for (const auto& piece : split_by(value->second, ',')) - term.constants.push_back(parse_double(piece, "constant value")); - term.dim = term.constants.size(); - std::vector names; - for (const double item : term.constants) - names.push_back(format_number(item)); - term.id = "constant:" + join(names, "+"); - validate_tokens_(tokens, {"type", "value", "index", "scale", "clip", "name"}, spec); - return finish_common_(term, tokens, spec); - } - - const auto path = tokens.find("path"); - if (path == tokens.end() || path->second.empty()) - throw std::invalid_argument( - "RlBridge: term requires path=... (term '" + spec - + "') unless it is joint_*/last_action/constant"); - term.kind = TermKind::kPath; - term.path = path->second; - const std::string take = string_token_(tokens, "take", ""); - const std::string transform = string_token_(tokens, "transform", ""); - - if (transform == "projected_gravity") { - if (!take.empty() && take != "gravity") - throw std::invalid_argument( - "RlBridge: transform=projected_gravity conflicts with " - "take=" - + take + " (term '" + spec + "')"); - term.take = Take::kGravity; - term.id = "gravity:" + term.path; - term.dim = 3; - if (!type.empty()) { - if (type != "quaternion") - throw std::invalid_argument( - "RlBridge: projected_gravity requires a " - "quaternion interface; type=" - + type + " is not quaternion (term '" + spec + "')"); - term.has_binding = true; - term.binding = Binding::kQuaternion; - } - } else if (transform == "gravity") { - throw std::invalid_argument( - "RlBridge: transform=gravity is not supported; use " - "transform=projected_gravity " - "(term '" - + spec + "')"); - } else if (!transform.empty()) { - throw std::invalid_argument( - "RlBridge: unknown transform='" + transform + "' (term '" + spec + "')"); - } else if (take.empty()) { - term.take = Take::kScalar; - term.id = term.path; - term.dim = 1; - } else if (take == "x" || take == "y" || take == "z") { - term.take = Take::kComponent; - term.component = std::string{"xyz"}.find(take); - term.id = "vec3c:" + term.path + ":" + take; - term.dim = 1; - } else if (take == "vec3" || take == "vec" || take == "all" || take == "vector") { - term.take = Take::kVector; - term.id = "vec3:" + term.path; - term.dim = 3; - } else if (take == "gravity") { - term.take = Take::kGravity; - term.id = "gravity:" + term.path; - term.dim = 3; + std::vector stacked; + stacked.reserve(obs_size_); + for (const auto& frame : history_) + stacked.insert(stacked.end(), frame.begin(), frame.end()); + publish_observation_(stacked, now); + ++pub_ok_count_; } else { - throw std::invalid_argument( - "RlBridge: unknown take='" + take - + "' (use x|y|z|vec3|gravity, or omit for scalar) (term '" + spec + "')"); - } - - if (!type.empty()) { - term.has_binding = true; - if (type == "scalar" || type == "double") { - term.binding = Binding::kDouble; - } else if (type == "bool") { - term.binding = Binding::kBool; - } else if (type == "int") { - term.binding = Binding::kInt; - } else if (type == "size" || type == "size_t") { - term.binding = Binding::kSize; - } else if (type == "vector3" || type == "vec3") { - term.binding = Binding::kVector3; - } else if (type == "direction_vector") { - term.binding = Binding::kDirectionVector; - } else if (type == "quaternion") { - term.binding = Binding::kQuaternion; - } else { - throw std::invalid_argument( - "RlBridge: unknown type='" + type + "' (term '" + spec + "')"); - } - const bool scalar_binding = - term.binding == Binding::kDouble || term.binding == Binding::kBool - || term.binding == Binding::kInt || term.binding == Binding::kSize; - const bool vector_binding = - term.binding == Binding::kVector3 || term.binding == Binding::kDirectionVector; - if (term.take == Take::kScalar && !scalar_binding) - throw std::invalid_argument( - "RlBridge: type=" + type - + " needs take=vec3 or take=x|y|z (a whole vector is not a scalar) (term '" - + spec + "')"); - if ((term.take == Take::kComponent || term.take == Take::kVector) && !vector_binding) - throw std::invalid_argument( - "RlBridge: take=" + take - + " needs a vector interface " - "(type=vector3|direction_vector), got type=" - + type + " (term '" + spec + "')"); - if (term.take == Take::kGravity && term.binding != Binding::kQuaternion) - throw std::invalid_argument( - "RlBridge: transform=projected_gravity needs " - "type=quaternion (term '" - + spec + "')"); + ++obs_invalid_count_; + RCLCPP_WARN_THROTTLE( + get_logger(), *get_clock(), 1000, + "observation sources not ready or non-finite; no obs published (%llu frames " + "skipped)", + static_cast(obs_invalid_count_)); } - - validate_tokens_( - tokens, - {"path", "take", "transform", "type", "index", "scale", "clip", "default", "name"}, - spec); - return finish_common_(term, tokens, spec); } - ObsTerm finish_common_( - ObsTerm term, const std::map& tokens, const std::string& spec) { - parse_index_(tokens, spec, term.has_index, term.index); - term.scale = double_token_(tokens, "scale", 1.0, spec); - if (term.scale == 0.0) - throw std::invalid_argument("RlBridge: scale=0 is not allowed (term '" + spec + "')"); - parse_clip_(tokens, spec, term.has_clip, term.clip_min, term.clip_max); - - const auto default_value = tokens.find("default"); - if (default_value != tokens.end()) { - if (!(term.kind == TermKind::kPath && term.take == Take::kScalar)) - throw std::invalid_argument( - "RlBridge: default= is only valid for scalar path " - "terms " - "(term '" - + spec + "')"); - term.default_value = parse_double(default_value->second, "default"); - if (!is_finite(term.default_value)) - throw std::invalid_argument( - "RlBridge: default must be finite (term '" + spec + "')"); - term.has_default = true; - } - - const auto name = tokens.find("name"); - if (name != tokens.end()) { - if (name->second.empty()) - throw std::invalid_argument( - "RlBridge: name= must not be empty (term '" + spec + "')"); - term.id = name->second; - } - return term; - } - - ActionTerm parse_action_term_(const std::string& spec) { - ActionTerm term; - const auto tokens = parse_tokens_(spec); - const auto index = tokens.find("index"); - if (index == tokens.end()) - throw std::invalid_argument( - "RlBridge: action term requires index= (term '" + spec + "')"); - const auto index_value = parse_integer(index->second, "action index"); - if (index_value < 0) - throw std::invalid_argument( - "RlBridge: action index must be >= 0 (term '" + spec + "')"); - term.index = static_cast(index_value); - - const auto output = tokens.find("output"); - if (output == tokens.end() || output->second.empty()) - throw std::invalid_argument( - "RlBridge: action term requires output= (term '" + spec + "')"); - term.output = output->second; - term.id = string_token_(tokens, "name", term.output); - if (term.id.empty()) - throw std::invalid_argument( - "RlBridge: action name= must not be empty (term '" + spec + "')"); - term.scale = double_token_(tokens, "scale", 1.0, spec); - if (term.scale == 0.0) - throw std::invalid_argument( - "RlBridge: action scale=0 is not allowed (term '" + spec + "')"); - parse_clip_(tokens, spec, term.has_clip, term.clip_min, term.clip_max); - validate_tokens_(tokens, {"index", "output", "name", "scale", "clip"}, spec); - return term; - } - - std::string obs_layout_signature_() const { - std::string signature = "v3-history=" + std::to_string(history_length_); - for (const auto& term : obs_terms_) { - signature += "|" + term.id; - if (term.scale != 1.0) - signature += "*" + format_number(term.scale); - signature += "@" + std::to_string(term.dim); - } - return signature; + bool read_enable_() const { + if (enable_slot_ == kNoSlot) + return enable_default_; + double raw = 0.0; + if (!read_double(slots_[enable_slot_], raw)) + return enable_default_; + return raw != 0.0; } - std::string actions_layout_signature_() const { - std::string signature = "v2"; - for (const auto& term : action_terms_) { - signature += "|#" + std::to_string(term.index) + ":" + term.id; - if (term.scale != 1.0) - signature += "*" + format_number(term.scale); + void check_contract_(const ActionSnapshot& snapshot) { + if (snapshot.layout_hash != layout_hash_) { + if (contract_ok_) { + contract_ok_ = false; + RCLCPP_FATAL( + get_logger(), + "contract mismatch: action layout_hash=%s != bridge layout_hash=%s " + "(policy process and bridge disagree; refusing to output actions)", + hex16(snapshot.layout_hash).c_str(), hex16(layout_hash_).c_str()); + } + } else if (expected_model_id_ != 0 && snapshot.model_id != expected_model_id_) { + if (contract_ok_) { + contract_ok_ = false; + RCLCPP_FATAL( + get_logger(), "model mismatch: action model_id=%s != expected_model_id=%s", + hex16(snapshot.model_id).c_str(), hex16(expected_model_id_).c_str()); + } } - return signature; } void log_layout_() const { @@ -1210,7 +395,7 @@ class RlBridge final if (!term.path.empty()) { line << " <- " << term.path; if (term.slot != kNoSlot) - line << " (" << binding_name_(slots_[term.slot]->binding) << ")"; + line << " (" << binding_name(slots_[term.slot].binding) << ")"; else line << " (default)"; } @@ -1226,175 +411,6 @@ class RlBridge final } } - bool read_double_(std::size_t slot, double& value) const { - const auto& entry = *slots_[slot]; - switch (entry.binding) { - case Binding::kDouble: - if (!entry.double_value->ready()) - return false; - value = **entry.double_value; - return true; - case Binding::kBool: - if (!entry.bool_value->ready()) - return false; - value = **entry.bool_value ? 1.0 : 0.0; - return true; - case Binding::kInt: - if (!entry.int_value->ready()) - return false; - value = static_cast(**entry.int_value); - return true; - case Binding::kSize: - if (!entry.size_value->ready()) - return false; - value = static_cast(**entry.size_value); - return true; - default: return false; - } - } - - bool read_unsigned_(std::size_t slot, std::uint64_t& value) const { - double raw = 0.0; - if (!read_double_(slot, raw)) - return false; - if (!is_finite(raw) || raw < 0.0) - return false; - value = static_cast(raw); - return true; - } - - bool read_enable_() const { - if (enable_slot_ == kNoSlot) - return enable_default_; - double raw = 0.0; - if (!read_double_(enable_slot_, raw)) - return enable_default_; - return raw != 0.0; - } - - bool build_observation_(std::vector& obs) const { - obs.assign(obs_size_, 0.0); - const auto push = [&obs]( - std::size_t index, double value, double scale, bool has_clip, - double clip_min, double clip_max, bool& ok) { - double result = value * scale; - if (has_clip) - result = std::clamp(result, clip_min, clip_max); - if (!is_finite(result)) - ok = false; - obs[index] = result; - }; - - for (const auto& term : obs_terms_) { - bool ok = true; - switch (term.kind) { - case TermKind::kPath: { - if (term.slot == kNoSlot) { - push( - term.index, term.default_value, term.scale, term.has_clip, term.clip_min, - term.clip_max, ok); - break; - } - const auto& entry = *slots_[term.slot]; - if (term.take == Take::kScalar) { - double value = 0.0; - if (!read_double_(term.slot, value)) - return false; - push( - term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, - ok); - } else if (term.take == Take::kComponent) { - double value = 0.0; - if (entry.binding == Binding::kVector3) { - if (!entry.vector3_value->ready()) - return false; - value = (**entry.vector3_value)[term.component]; - } else { - if (!entry.direction_vector_value->ready()) - return false; - value = (**entry.direction_vector_value).vector[term.component]; - } - push( - term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, - ok); - } else if (term.take == Take::kVector) { - if (entry.binding == Binding::kVector3) { - if (!entry.vector3_value->ready()) - return false; - const auto& value = **entry.vector3_value; - for (std::size_t i = 0; i < 3; ++i) - push( - term.index + i, value[static_cast(i)], term.scale, - term.has_clip, term.clip_min, term.clip_max, ok); - } else { - if (!entry.direction_vector_value->ready()) - return false; - const auto& value = (**entry.direction_vector_value).vector; - for (std::size_t i = 0; i < 3; ++i) - push( - term.index + i, value[static_cast(i)], term.scale, - term.has_clip, term.clip_min, term.clip_max, ok); - } - } else { - if (!entry.quaternion_value->ready()) - return false; - const Eigen::Vector3d gravity = - **entry.quaternion_value * Eigen::Vector3d(0.0, 0.0, -1.0); - for (std::size_t i = 0; i < 3; ++i) - push( - term.index + i, gravity[static_cast(i)], term.scale, - term.has_clip, term.clip_min, term.clip_max, ok); - } - break; - } - case TermKind::kJointPos: { - for (std::size_t j = 0; j < term.joints.size(); ++j) { - double value = 0.0; - if (!term.zero) { - if (!read_double_(term.joint_slots[j], value)) - return false; - if (term.relative) - value -= term.joint_defaults[j]; - } - push( - term.index + j, value, term.scale, term.has_clip, term.clip_min, - term.clip_max, ok); - } - break; - } - case TermKind::kJointVel: - case TermKind::kJointTorque: { - for (std::size_t j = 0; j < term.joints.size(); ++j) { - double value = 0.0; - if (!read_double_(term.joint_slots[j], value)) - return false; - push( - term.index + j, value, term.scale, term.has_clip, term.clip_min, - term.clip_max, ok); - } - break; - } - case TermKind::kLastAction: { - for (std::size_t j = 0; j < term.action_indices.size(); ++j) - push( - term.index + j, last_actions_[term.action_indices[j]], term.scale, - term.has_clip, term.clip_min, term.clip_max, ok); - break; - } - case TermKind::kConstant: { - for (std::size_t j = 0; j < term.constants.size(); ++j) - push( - term.index + j, term.constants[j], term.scale, term.has_clip, term.clip_min, - term.clip_max, ok); - break; - } - } - if (!ok) - return false; - } - return true; - } - void publish_observation_( const std::vector& obs, std::chrono::steady_clock::time_point now) { rmcs_rl::msg::Observation message; @@ -1412,43 +428,16 @@ class RlBridge final } void on_action_(rmcs_rl::msg::Action::UniquePtr message) { - const auto received = std::chrono::steady_clock::now(); if (message->action.size() != action_size_) { RCLCPP_ERROR_THROTTLE( get_logger(), *get_clock(), 1000, "ignoring action with %zu values (expected %zu)", message->action.size(), action_size_); return; } - - const std::uint64_t sequence = action_sequence_.load(std::memory_order_relaxed); - action_sequence_.store(sequence + 1, std::memory_order_relaxed); - std::atomic_thread_fence(std::memory_order_release); - - incoming_.obs_seq = message->obs_seq; - incoming_.layout_hash = message->layout_hash; - incoming_.model_id = message->model_id; - incoming_.received = received; - std::copy(message->action.begin(), message->action.end(), incoming_.action.begin()); - - std::atomic_thread_fence(std::memory_order_release); - action_sequence_.store(sequence + 2, std::memory_order_relaxed); - } - - bool try_read_action_(ActionSnapshot& snapshot) { - for (int attempt = 0; attempt < 4; ++attempt) { - const std::uint64_t before = action_sequence_.load(std::memory_order_relaxed); - if (before == 0 || (before & 1U) != 0) - return false; - std::atomic_thread_fence(std::memory_order_acquire); - snapshot = incoming_; - std::atomic_thread_fence(std::memory_order_acquire); - if (action_sequence_.load(std::memory_order_relaxed) == before) - return true; - } - return false; + action_channel_.store(*message); } - double obs_age_of_(std::uint64_t obs_seq, std::chrono::steady_clock::time_point now) const { + double obs_age_of(std::uint64_t obs_seq, std::chrono::steady_clock::time_point now) const { if (!pub_started_) return std::numeric_limits::quiet_NaN(); if (obs_seq == pub_seq_) @@ -1458,62 +447,18 @@ class RlBridge final return std::numeric_limits::quiet_NaN(); } - void write_actions_(bool valid, const ActionSnapshot& snapshot) { - for (std::size_t i = 0; i < action_terms_.size(); ++i) { - const auto& term = action_terms_[i]; - double value = 0.0; - if (valid) { - value = snapshot.action[i] * term.scale; - if (term.has_clip) - value = std::clamp(value, term.clip_min, term.clip_max); - written_[i] = value; - } else { - switch (invalid_mode_) { - case InvalidMode::kNaN: value = std::numeric_limits::quiet_NaN(); break; - case InvalidMode::kZero: value = 0.0; break; - case InvalidMode::kHold: value = written_[i]; break; - } - } - **action_outputs_[i] = value; - } - } - void reset_runtime_() { - std::fill(last_actions_.begin(), last_actions_.end(), 0.0); - std::fill(written_.begin(), written_.end(), 0.0); + reset_action_state(last_actions_, written_); history_.clear(); pub_started_ = false; prev_pub_seq_ = 0; } - std::string invalid_reason_( - bool enabled, bool contract_ok, bool has_snapshot, bool fresh, bool seq_ok, - bool finite) const { - if (!contract_ok) - return "contract/model mismatch (latched; restart or re-enable)"; - if (!enabled) - return "disabled by enable interface"; - if (!has_snapshot) - return "no action received yet"; - if (!finite) - return "action contains non-finite values"; - if (!seq_ok) - return "action answers an obs frame older than the previous one"; - if (!fresh) - return "action is stale (age > max_action_age)"; - return "unknown"; - } - - std::vector joint_names_; - std::string joint_base_path_; - std::unordered_map joint_index_; - std::unordered_map angle_suffix_; - std::unordered_map velocity_suffix_; - std::unordered_map torque_suffix_; + JointConfig joint_config_; std::vector obs_terms_; std::vector action_terms_; - std::vector> slots_; + std::vector slots_; std::vector>> action_outputs_; std::unordered_set own_output_paths_; @@ -1553,9 +498,8 @@ class RlBridge final bool contract_ok_ = true; bool last_valid_ = false; - ActionSnapshot incoming_; + ActionChannel action_channel_; ActionSnapshot read_snapshot_; - alignas(64) std::atomic action_sequence_{0}; OutputInterface valid_output_{}; OutputInterface healthy_output_{}; diff --git a/src/rl_bridge/action_channel.cpp b/src/rl_bridge/action_channel.cpp new file mode 100644 index 0000000..7237f54 --- /dev/null +++ b/src/rl_bridge/action_channel.cpp @@ -0,0 +1,55 @@ +#include "rl_bridge/action_channel.hpp" + +#include +#include + +namespace rmcs_rl { + +void write_actions( + bool valid, const ActionSnapshot& snapshot, const std::vector& terms, + InvalidMode invalid_mode, std::vector& written, + std::vector>>& outputs) { + for (std::size_t i = 0; i < terms.size(); ++i) { + const auto& term = terms[i]; + double value = 0.0; + if (valid) { + value = snapshot.action[i] * term.scale; + if (term.has_clip) + value = std::clamp(value, term.clip_min, term.clip_max); + written[i] = value; + } else { + switch (invalid_mode) { + case InvalidMode::kNaN: + value = std::numeric_limits::quiet_NaN(); + break; + case InvalidMode::kZero: value = 0.0; break; + case InvalidMode::kHold: value = written[i]; break; + } + } + **outputs[i] = value; + } +} + +void reset_action_state(std::vector& last_actions, std::vector& written) { + std::fill(last_actions.begin(), last_actions.end(), 0.0); + std::fill(written.begin(), written.end(), 0.0); +} + +std::string invalid_reason( + bool enabled, bool contract_ok, bool has_snapshot, bool fresh, bool seq_ok, bool finite) { + if (!contract_ok) + return "contract/model mismatch (latched; restart or re-enable)"; + if (!enabled) + return "disabled by enable interface"; + if (!has_snapshot) + return "no action received yet"; + if (!finite) + return "action contains non-finite values"; + if (!seq_ok) + return "action answers an obs frame older than the previous one"; + if (!fresh) + return "action is stale (age > max_action_age)"; + return "unknown"; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/action_channel.hpp b/src/rl_bridge/action_channel.hpp new file mode 100644 index 0000000..f7eca2d --- /dev/null +++ b/src/rl_bridge/action_channel.hpp @@ -0,0 +1,69 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include + +#include "rl_bridge/types.hpp" + +namespace rmcs_rl { + +class ActionChannel { +public: + void resize(std::size_t action_size) { incoming_.action.assign(action_size, 0.0); } + + void store(const rmcs_rl::msg::Action& message) { + const auto received = std::chrono::steady_clock::now(); + + const std::uint64_t sequence = action_sequence_.load(std::memory_order_relaxed); + action_sequence_.store(sequence + 1, std::memory_order_relaxed); + std::atomic_thread_fence(std::memory_order_release); + + incoming_.obs_seq = message.obs_seq; + incoming_.layout_hash = message.layout_hash; + incoming_.model_id = message.model_id; + incoming_.received = received; + std::copy(message.action.begin(), message.action.end(), incoming_.action.begin()); + + std::atomic_thread_fence(std::memory_order_release); + action_sequence_.store(sequence + 2, std::memory_order_relaxed); + } + + bool try_read(ActionSnapshot& snapshot) { + for (int attempt = 0; attempt < 4; ++attempt) { + const std::uint64_t before = action_sequence_.load(std::memory_order_relaxed); + if (before == 0 || (before & 1U) != 0) + return false; + std::atomic_thread_fence(std::memory_order_acquire); + snapshot = incoming_; + std::atomic_thread_fence(std::memory_order_acquire); + if (action_sequence_.load(std::memory_order_relaxed) == before) + return true; + } + return false; + } + +private: + ActionSnapshot incoming_; + alignas(64) std::atomic action_sequence_{0}; +}; + +void write_actions( + bool valid, const ActionSnapshot& snapshot, const std::vector& terms, + InvalidMode invalid_mode, std::vector& written, + std::vector>>& outputs); + +void reset_action_state(std::vector& last_actions, std::vector& written); + +std::string invalid_reason( + bool enabled, bool contract_ok, bool has_snapshot, bool fresh, bool seq_ok, bool finite); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/interface_binding.cpp b/src/rl_bridge/interface_binding.cpp new file mode 100644 index 0000000..819a0f2 --- /dev/null +++ b/src/rl_bridge/interface_binding.cpp @@ -0,0 +1,157 @@ +#include "rl_bridge/interface_binding.hpp" + +#include +#include + +#include "rl_bridge/utility.hpp" + +namespace rmcs_rl { + +const char* binding_name(Binding binding) { + switch (binding) { + case Binding::kDouble: return "double"; + case Binding::kBool: return "bool"; + case Binding::kInt: return "int"; + case Binding::kSize: return "std::size_t"; + case Binding::kVector3: return "Eigen::Vector3d"; + case Binding::kDirectionVector: return "BaseLink::DirectionVector"; + case Binding::kQuaternion: return "Eigen::Quaterniond"; + } + return "unknown"; +} + +bool binding_matches_type(Binding binding, const std::type_info& type) { + switch (binding) { + case Binding::kDouble: return type == typeid(double); + case Binding::kBool: return type == typeid(bool); + case Binding::kInt: return type == typeid(int); + case Binding::kSize: return type == typeid(std::size_t); + case Binding::kVector3: return type == typeid(Eigen::Vector3d); + case Binding::kDirectionVector: + return type == typeid(rmcs_description::BaseLink::DirectionVector); + case Binding::kQuaternion: return type == typeid(Eigen::Quaterniond); + } + return false; +} + +std::size_t acquire_slot( + rmcs_executor::Component& component, const std::string& path, + const std::vector& candidates, bool required, + const rmcs_executor::Component::OutputInfoMap& output_map, std::vector& slots, + const char* context) { + const auto output = output_map.find(path); + if (output == output_map.end()) { + if (!required) + return kNoSlot; + throw std::runtime_error( + "RlBridge: required input interface \"" + path + + "\" was not produced by any component (" + context + + "); check the interface path or add default= to make it optional"); + } + if (output->second.kind != rmcs_executor::InterfaceKind::Normal) + throw std::runtime_error( + "RlBridge: input interface \"" + path + + "\" exists but is an Event interface; Normal required (" + context + ")"); + const std::type_info& producer_type = output->second.type.get(); + + for (std::size_t i = 0; i < slots.size(); ++i) + if (slots[i].path == path && binding_matches_type(slots[i].binding, producer_type)) + return i; + + Binding selected = Binding::kDouble; + bool found = false; + for (const auto binding : candidates) + if (binding_matches_type(binding, producer_type)) { + selected = binding; + found = true; + break; + } + if (!found) { + std::vector expected; + for (const auto binding : candidates) + expected.push_back(binding_name(binding)); + throw std::runtime_error( + "RlBridge: cannot bind observation interface \"" + path + "\": producer declares \"" + + pretty_type(producer_type) + "\" but the term accepts { " + join(expected, ", ") + + " }. Either fix the term (take=/transform=) or pin type= explicitly."); + } + + Slot slot; + slot.path = path; + slot.binding = selected; + slot.required = required; + switch (selected) { + case Binding::kDouble: + slot.double_value = std::make_unique>(); + component.register_input(path, *slot.double_value, required); + break; + case Binding::kBool: + slot.bool_value = std::make_unique>(); + component.register_input(path, *slot.bool_value, required); + break; + case Binding::kInt: + slot.int_value = std::make_unique>(); + component.register_input(path, *slot.int_value, required); + break; + case Binding::kSize: + slot.size_value = + std::make_unique>(); + component.register_input(path, *slot.size_value, required); + break; + case Binding::kVector3: + slot.vector3_value = + std::make_unique>(); + component.register_input(path, *slot.vector3_value, required); + break; + case Binding::kDirectionVector: + slot.direction_vector_value = std::make_unique< + rmcs_executor::Component::InputInterface>(); + component.register_input(path, *slot.direction_vector_value, required); + break; + case Binding::kQuaternion: + slot.quaternion_value = + std::make_unique>(); + component.register_input(path, *slot.quaternion_value, required); + break; + } + slots.push_back(std::move(slot)); + return slots.size() - 1; +} + +bool read_double(const Slot& slot, double& value) { + switch (slot.binding) { + case Binding::kDouble: + if (!slot.double_value->ready()) + return false; + value = **slot.double_value; + return true; + case Binding::kBool: + if (!slot.bool_value->ready()) + return false; + value = **slot.bool_value ? 1.0 : 0.0; + return true; + case Binding::kInt: + if (!slot.int_value->ready()) + return false; + value = static_cast(**slot.int_value); + return true; + case Binding::kSize: + if (!slot.size_value->ready()) + return false; + value = static_cast(**slot.size_value); + return true; + default: return false; + } +} + +bool read_unsigned(const Slot& slot, std::uint64_t& value) { + double raw = 0.0; + if (!read_double(slot, raw)) + return false; + if (!std::isfinite(raw) || raw < 0.0) + return false; + value = static_cast(raw); + return true; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/interface_binding.hpp b/src/rl_bridge/interface_binding.hpp new file mode 100644 index 0000000..9224374 --- /dev/null +++ b/src/rl_bridge/interface_binding.hpp @@ -0,0 +1,26 @@ +#pragma once + +#include +#include +#include +#include + +#include + +#include "rl_bridge/types.hpp" + +namespace rmcs_rl { + +const char* binding_name(Binding binding); +bool binding_matches_type(Binding binding, const std::type_info& type); + +std::size_t acquire_slot( + rmcs_executor::Component& component, const std::string& path, + const std::vector& candidates, bool required, + const rmcs_executor::Component::OutputInfoMap& output_map, std::vector& slots, + const char* context); + +bool read_double(const Slot& slot, double& value); +bool read_unsigned(const Slot& slot, std::uint64_t& value); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/joint_config.cpp b/src/rl_bridge/joint_config.cpp new file mode 100644 index 0000000..bb8c4a9 --- /dev/null +++ b/src/rl_bridge/joint_config.cpp @@ -0,0 +1,81 @@ +#include "rl_bridge/joint_config.hpp" + +#include + +#include "rl_bridge/parameters.hpp" +#include "rl_bridge/utility.hpp" + +namespace rmcs_rl { + +JointConfig load_joint_config(rclcpp::Node& node) { + JointConfig config; + config.names = string_array_or(node, "joint_names"); + config.base_path = string_or(node, "joint_base_path", ""); + if (!config.names.empty() && config.base_path.empty()) + throw std::invalid_argument("RlBridge: joint_base_path is required when joint_names is set"); + for (std::size_t i = 0; i < config.names.size(); ++i) { + if (config.names[i].empty()) + throw std::invalid_argument("RlBridge: joint_names contains an empty name"); + if (!config.index.emplace(config.names[i], i).second) + throw std::invalid_argument( + "RlBridge: duplicate joint name '" + config.names[i] + "'"); + } + + const std::string default_angle = string_or(node, "joint_angle_suffix", "/angle"); + const std::string default_velocity = string_or(node, "joint_velocity_suffix", "/velocity"); + const std::string default_torque = string_or(node, "joint_torque_suffix", "/torque"); + for (const auto& joint : config.names) { + config.angle_suffix[joint] = default_angle; + config.velocity_suffix[joint] = default_velocity; + config.torque_suffix[joint] = default_torque; + } + + constexpr const char* prefix = "joint_interface_overrides."; + for (const auto& name : node.list_parameters({"joint_interface_overrides"}, 10).names) { + if (name.rfind(prefix, 0) != 0) + continue; + const std::string remainder = name.substr(std::string{prefix}.size()); + const auto separator = remainder.find_last_of('.'); + if (separator == std::string::npos) + throw std::invalid_argument( + "RlBridge: joint_interface_overrides entry '" + name + + "' must be ."); + const std::string joint = remainder.substr(0, separator); + const std::string field = remainder.substr(separator + 1); + if (config.index.count(joint) == 0) + throw std::invalid_argument( + "RlBridge: joint_interface_overrides references unknown joint '" + joint + "'"); + const std::string value = string_or(node, name, ""); + if (field == "angle") + config.angle_suffix[joint] = value; + else if (field == "velocity") + config.velocity_suffix[joint] = value; + else if (field == "torque") + config.torque_suffix[joint] = value; + else + throw std::invalid_argument( + "RlBridge: joint_interface_overrides field '" + field + + "' is not angle/velocity/torque"); + } + return config; +} + +std::string joint_path(const JointConfig& config, const std::string& joint, char field) { + const auto suffix = field == 'a' ? config.angle_suffix.at(joint) + : field == 'v' ? config.velocity_suffix.at(joint) + : config.torque_suffix.at(joint); + return config.base_path + "/" + joint + suffix; +} + +std::optional +default_joint_pos(rclcpp::Node& node, const JointConfig& config, std::size_t joint) { + const auto name = "default_joint_pos." + config.names[joint]; + const auto value = number_parameter(node, name); + if (!value.has_value()) + return std::nullopt; + if (!is_finite(*value)) + throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); + return value; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/joint_config.hpp b/src/rl_bridge/joint_config.hpp new file mode 100644 index 0000000..8c7d3f8 --- /dev/null +++ b/src/rl_bridge/joint_config.hpp @@ -0,0 +1,27 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include + +namespace rmcs_rl { + +struct JointConfig { + std::vector names; + std::string base_path; + std::unordered_map index; + std::unordered_map angle_suffix; + std::unordered_map velocity_suffix; + std::unordered_map torque_suffix; +}; + +JointConfig load_joint_config(rclcpp::Node& node); +std::string joint_path(const JointConfig& config, const std::string& joint, char field); +std::optional +default_joint_pos(rclcpp::Node& node, const JointConfig& config, std::size_t joint); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/observation.cpp b/src/rl_bridge/observation.cpp new file mode 100644 index 0000000..f11e384 --- /dev/null +++ b/src/rl_bridge/observation.cpp @@ -0,0 +1,156 @@ +#include "rl_bridge/observation.hpp" + +#include +#include +#include + +#include "rl_bridge/interface_binding.hpp" +#include "rl_bridge/utility.hpp" + +namespace rmcs_rl { + +bool build_observation( + const std::vector& terms, const std::vector& slots, + const std::vector& last_actions, std::size_t obs_size, std::vector& obs) { + obs.assign(obs_size, 0.0); + const auto push = [&obs]( + std::size_t index, double value, double scale, bool has_clip, + double clip_min, double clip_max, bool& ok) { + double result = value * scale; + if (has_clip) + result = std::clamp(result, clip_min, clip_max); + if (!is_finite(result)) + ok = false; + obs[index] = result; + }; + + for (const auto& term : terms) { + bool ok = true; + switch (term.kind) { + case TermKind::kPath: { + if (term.slot == kNoSlot) { + push( + term.index, term.default_value, term.scale, term.has_clip, term.clip_min, + term.clip_max, ok); + break; + } + const auto& entry = slots[term.slot]; + if (term.take == Take::kScalar) { + double value = 0.0; + if (!read_double(entry, value)) + return false; + push( + term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); + } else if (term.take == Take::kComponent) { + double value = 0.0; + if (entry.binding == Binding::kVector3) { + if (!entry.vector3_value->ready()) + return false; + value = (**entry.vector3_value)[term.component]; + } else { + if (!entry.direction_vector_value->ready()) + return false; + value = (**entry.direction_vector_value).vector[term.component]; + } + push( + term.index, value, term.scale, term.has_clip, term.clip_min, term.clip_max, ok); + } else if (term.take == Take::kVector) { + if (entry.binding == Binding::kVector3) { + if (!entry.vector3_value->ready()) + return false; + const auto& value = **entry.vector3_value; + for (std::size_t i = 0; i < 3; ++i) + push( + term.index + i, value[static_cast(i)], term.scale, + term.has_clip, term.clip_min, term.clip_max, ok); + } else { + if (!entry.direction_vector_value->ready()) + return false; + const auto& value = (**entry.direction_vector_value).vector; + for (std::size_t i = 0; i < 3; ++i) + push( + term.index + i, value[static_cast(i)], term.scale, + term.has_clip, term.clip_min, term.clip_max, ok); + } + } else { + if (!entry.quaternion_value->ready()) + return false; + const Eigen::Vector3d gravity = + **entry.quaternion_value * Eigen::Vector3d(0.0, 0.0, -1.0); + for (std::size_t i = 0; i < 3; ++i) + push( + term.index + i, gravity[static_cast(i)], term.scale, + term.has_clip, term.clip_min, term.clip_max, ok); + } + break; + } + case TermKind::kJointPos: { + for (std::size_t j = 0; j < term.joints.size(); ++j) { + double value = 0.0; + if (!term.zero) { + if (!read_double(slots[term.joint_slots[j]], value)) + return false; + if (term.relative) + value -= term.joint_defaults[j]; + } + push( + term.index + j, value, term.scale, term.has_clip, term.clip_min, term.clip_max, + ok); + } + break; + } + case TermKind::kJointVel: + case TermKind::kJointTorque: { + for (std::size_t j = 0; j < term.joints.size(); ++j) { + double value = 0.0; + if (!read_double(slots[term.joint_slots[j]], value)) + return false; + push( + term.index + j, value, term.scale, term.has_clip, term.clip_min, term.clip_max, + ok); + } + break; + } + case TermKind::kLastAction: { + for (std::size_t j = 0; j < term.action_indices.size(); ++j) + push( + term.index + j, last_actions[term.action_indices[j]], term.scale, term.has_clip, + term.clip_min, term.clip_max, ok); + break; + } + case TermKind::kConstant: { + for (std::size_t j = 0; j < term.constants.size(); ++j) + push( + term.index + j, term.constants[j], term.scale, term.has_clip, term.clip_min, + term.clip_max, ok); + break; + } + } + if (!ok) + return false; + } + return true; +} + +std::string obs_layout_signature(const std::vector& terms, std::size_t history_length) { + std::string signature = "v3-history=" + std::to_string(history_length); + for (const auto& term : terms) { + signature += "|" + term.id; + if (term.scale != 1.0) + signature += "*" + format_number(term.scale); + signature += "@" + std::to_string(term.dim); + } + return signature; +} + +std::string actions_layout_signature(const std::vector& terms) { + std::string signature = "v2"; + for (const auto& term : terms) { + signature += "|#" + std::to_string(term.index) + ":" + term.id; + if (term.scale != 1.0) + signature += "*" + format_number(term.scale); + } + return signature; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/observation.hpp b/src/rl_bridge/observation.hpp new file mode 100644 index 0000000..f663542 --- /dev/null +++ b/src/rl_bridge/observation.hpp @@ -0,0 +1,18 @@ +#pragma once + +#include +#include +#include + +#include "rl_bridge/types.hpp" + +namespace rmcs_rl { + +bool build_observation( + const std::vector& terms, const std::vector& slots, + const std::vector& last_actions, std::size_t obs_size, std::vector& obs); + +std::string obs_layout_signature(const std::vector& terms, std::size_t history_length); +std::string actions_layout_signature(const std::vector& terms); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/parameters.cpp b/src/rl_bridge/parameters.cpp new file mode 100644 index 0000000..04dca2c --- /dev/null +++ b/src/rl_bridge/parameters.cpp @@ -0,0 +1,99 @@ +#include "rl_bridge/parameters.hpp" + +#include +#include + +namespace rmcs_rl { + +std::optional number_parameter(rclcpp::Node& node, const std::string& name) { + rclcpp::Parameter parameter; + try { + if (!node.get_parameter(name, parameter)) + return std::nullopt; + } catch (const std::exception&) { + return std::nullopt; + } + switch (parameter.get_type()) { + case rclcpp::ParameterType::PARAMETER_DOUBLE: return parameter.as_double(); + case rclcpp::ParameterType::PARAMETER_INTEGER: return static_cast(parameter.as_int()); + case rclcpp::ParameterType::PARAMETER_NOT_SET: return std::nullopt; + default: throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a number"); + } +} + +double number_or(rclcpp::Node& node, const std::string& name, double fallback) { + const auto value = number_parameter(node, name); + if (!value.has_value()) + return fallback; + if (!std::isfinite(*value)) + throw std::invalid_argument("RlBridge: parameter '" + name + "' is not finite"); + return *value; +} + +std::uint64_t parse_u64(const std::string& text, const std::string& name) { + try { + std::size_t consumed = 0; + const bool hex = text.rfind("0x", 0) == 0 || text.rfind("0X", 0) == 0; + const auto value = std::stoull(hex ? text.substr(2) : text, &consumed, hex ? 16 : 10); + if (consumed != (hex ? text.size() - 2 : text.size())) + throw std::invalid_argument("trailing characters"); + return static_cast(value); + } catch (const std::exception&) { + throw std::invalid_argument( + "RlBridge: parameter '" + name + + "' must be a decimal or 0x-prefixed 64-bit id, got '" + text + "'"); + } +} + +std::optional integer_parameter(rclcpp::Node& node, const std::string& name) { + const auto value = number_parameter(node, name); + if (!value.has_value()) + return std::nullopt; + const auto rounded = std::llround(*value); + if (std::abs(*value - static_cast(rounded)) > 1e-9) + throw std::invalid_argument("RlBridge: parameter '" + name + "' must be an integer"); + return rounded; +} + +bool bool_or(rclcpp::Node& node, const std::string& name, bool fallback) { + rclcpp::Parameter parameter; + try { + if (!node.get_parameter(name, parameter)) + return fallback; + } catch (const std::exception&) { + return fallback; + } + if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) + return fallback; + if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_BOOL) + throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a boolean"); + return parameter.as_bool(); +} + +std::string string_or(rclcpp::Node& node, const std::string& name, const std::string& fallback) { + rclcpp::Parameter parameter; + try { + if (!node.get_parameter(name, parameter)) + return fallback; + } catch (const std::exception&) { + return fallback; + } + if (parameter.get_type() == rclcpp::ParameterType::PARAMETER_NOT_SET) + return fallback; + if (parameter.get_type() != rclcpp::ParameterType::PARAMETER_STRING) + throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a string"); + return parameter.as_string(); +} + +std::vector string_array_or(rclcpp::Node& node, const std::string& name) { + std::vector value; + try { + if (!node.get_parameter(name, value)) + return {}; + } catch (const std::exception&) { + throw std::invalid_argument("RlBridge: parameter '" + name + "' must be a list of strings"); + } + return value; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/parameters.hpp b/src/rl_bridge/parameters.hpp new file mode 100644 index 0000000..0d43c41 --- /dev/null +++ b/src/rl_bridge/parameters.hpp @@ -0,0 +1,20 @@ +#pragma once + +#include +#include +#include +#include + +#include + +namespace rmcs_rl { + +std::optional number_parameter(rclcpp::Node& node, const std::string& name); +double number_or(rclcpp::Node& node, const std::string& name, double fallback); +std::uint64_t parse_u64(const std::string& text, const std::string& name); +std::optional integer_parameter(rclcpp::Node& node, const std::string& name); +bool bool_or(rclcpp::Node& node, const std::string& name, bool fallback); +std::string string_or(rclcpp::Node& node, const std::string& name, const std::string& fallback); +std::vector string_array_or(rclcpp::Node& node, const std::string& name); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/term_parser.cpp b/src/rl_bridge/term_parser.cpp new file mode 100644 index 0000000..1db287d --- /dev/null +++ b/src/rl_bridge/term_parser.cpp @@ -0,0 +1,285 @@ +#include "rl_bridge/term_parser.hpp" + +#include +#include + +#include "rl_bridge/utility.hpp" + +namespace rmcs_rl { + +namespace { + +ObsTerm finish_common( + ObsTerm term, const std::map& tokens, const std::string& spec) { + parse_index(tokens, spec, term.has_index, term.index); + term.scale = double_token(tokens, "scale", 1.0, spec); + if (term.scale == 0.0) + throw std::invalid_argument("RlBridge: scale=0 is not allowed (term '" + spec + "')"); + parse_clip(tokens, spec, term.has_clip, term.clip_min, term.clip_max); + + const auto default_value = tokens.find("default"); + if (default_value != tokens.end()) { + if (!(term.kind == TermKind::kPath && term.take == Take::kScalar)) + throw std::invalid_argument( + "RlBridge: default= is only valid for scalar path " + "terms " + "(term '" + + spec + "')"); + term.default_value = parse_double(default_value->second, "default"); + if (!is_finite(term.default_value)) + throw std::invalid_argument( + "RlBridge: default must be finite (term '" + spec + "')"); + term.has_default = true; + } + + const auto name = tokens.find("name"); + if (name != tokens.end()) { + if (name->second.empty()) + throw std::invalid_argument("RlBridge: name= must not be empty (term '" + spec + "')"); + term.id = name->second; + } + return term; +} + +} // namespace + +ObsTerm parse_obs_term(const std::string& spec, const TermParseContext& context) { + if (spec.empty()) + throw std::invalid_argument("RlBridge: observation term must not be empty"); + const auto tokens = parse_tokens(spec); + ObsTerm term; + + std::string type = string_token(tokens, "type", ""); + + if (type == "joint_pos" || type == "joint_vel" || type == "joint_torque") { + term.kind = type == "joint_pos" ? TermKind::kJointPos + : type == "joint_vel" ? TermKind::kJointVel + : TermKind::kJointTorque; + const auto joints = tokens.find("joints"); + if (joints == tokens.end() || joints->second.empty()) + throw std::invalid_argument( + "RlBridge: joint_* term requires joints=a,b (term '" + spec + "')"); + for (const auto& name : split_by(joints->second, ',')) { + if (name.empty()) + throw std::invalid_argument( + "RlBridge: empty joint name in joints= (term '" + spec + "')"); + const auto iter = context.joints.index.find(name); + if (iter == context.joints.index.end()) + throw std::invalid_argument( + "RlBridge: unknown joint '" + name + "' (term '" + spec + "')"); + term.joints.push_back(iter->second); + term.joint_names.push_back(name); + } + term.relative = bool_token(tokens, "relative", false, spec); + term.zero = bool_token(tokens, "zero", false, spec); + if (term.kind != TermKind::kJointPos && (term.relative || term.zero)) + throw std::invalid_argument( + "RlBridge: relative/zero are only valid for joint_pos (term '" + spec + "')"); + if (term.relative && term.zero) + throw std::invalid_argument( + "RlBridge: joint_pos cannot be both relative and zero (term '" + spec + "')"); + for (const auto joint : term.joints) { + if (term.kind == TermKind::kJointPos && term.relative) { + const auto base = default_joint_pos(context.node, context.joints, joint); + if (!base.has_value()) + throw std::invalid_argument( + "RlBridge: joint_pos relative term '" + spec + + "' requires default_joint_pos." + context.joints.names[joint]); + term.joint_defaults.push_back(*base); + } else { + term.joint_defaults.push_back(0.0); + } + } + term.dim = term.joints.size(); + const std::string flag = term.zero ? "zero" : (term.relative ? "rel" : "abs"); + const std::string prefix = type == "joint_pos" ? "joint_pos" + : type == "joint_vel" ? "joint_vel" + : "joint_torque"; + term.id = prefix + ":" + flag + ":" + join(term.joint_names, "+"); + validate_tokens( + tokens, {"type", "joints", "relative", "zero", "index", "scale", "clip", "name"}, spec); + return finish_common(term, tokens, spec); + } + if (type == "last_action") { + term.kind = TermKind::kLastAction; + if (const auto indices = tokens.find("indices"); indices != tokens.end()) + term.action_indices = parse_index_list(indices->second, context.action_size, "last_action"); + else + for (std::size_t i = 0; i < context.action_size; ++i) + term.action_indices.push_back(i); + term.dim = term.action_indices.size(); + std::vector names; + for (const auto index : term.action_indices) + names.push_back(std::to_string(index)); + term.id = "last_action:" + join(names, "+"); + validate_tokens(tokens, {"type", "indices", "index", "scale", "clip", "name"}, spec); + return finish_common(term, tokens, spec); + } + if (type == "constant") { + term.kind = TermKind::kConstant; + const auto value = tokens.find("value"); + if (value == tokens.end() || value->second.empty()) + throw std::invalid_argument( + "RlBridge: constant term requires value=... (term '" + spec + "')"); + for (const auto& piece : split_by(value->second, ',')) + term.constants.push_back(parse_double(piece, "constant value")); + term.dim = term.constants.size(); + std::vector names; + for (const double item : term.constants) + names.push_back(format_number(item)); + term.id = "constant:" + join(names, "+"); + validate_tokens(tokens, {"type", "value", "index", "scale", "clip", "name"}, spec); + return finish_common(term, tokens, spec); + } + + const auto path = tokens.find("path"); + if (path == tokens.end() || path->second.empty()) + throw std::invalid_argument( + "RlBridge: term requires path=... (term '" + spec + + "') unless it is joint_*/last_action/constant"); + term.kind = TermKind::kPath; + term.path = path->second; + const std::string take = string_token(tokens, "take", ""); + const std::string transform = string_token(tokens, "transform", ""); + + if (transform == "projected_gravity") { + if (!take.empty() && take != "gravity") + throw std::invalid_argument( + "RlBridge: transform=projected_gravity conflicts with " + "take=" + + take + " (term '" + spec + "')"); + term.take = Take::kGravity; + term.id = "gravity:" + term.path; + term.dim = 3; + if (!type.empty()) { + if (type != "quaternion") + throw std::invalid_argument( + "RlBridge: projected_gravity requires a " + "quaternion interface; type=" + + type + " is not quaternion (term '" + spec + "')"); + term.has_binding = true; + term.binding = Binding::kQuaternion; + } + } else if (transform == "gravity") { + throw std::invalid_argument( + "RlBridge: transform=gravity is not supported; use " + "transform=projected_gravity " + "(term '" + + spec + "')"); + } else if (!transform.empty()) { + throw std::invalid_argument( + "RlBridge: unknown transform='" + transform + "' (term '" + spec + "')"); + } else if (take.empty()) { + term.take = Take::kScalar; + term.id = term.path; + term.dim = 1; + } else if (take == "x" || take == "y" || take == "z") { + term.take = Take::kComponent; + term.component = std::string{"xyz"}.find(take); + term.id = "vec3c:" + term.path + ":" + take; + term.dim = 1; + } else if (take == "vec3" || take == "vec" || take == "all" || take == "vector") { + term.take = Take::kVector; + term.id = "vec3:" + term.path; + term.dim = 3; + } else if (take == "gravity") { + term.take = Take::kGravity; + term.id = "gravity:" + term.path; + term.dim = 3; + } else { + throw std::invalid_argument( + "RlBridge: unknown take='" + take + + "' (use x|y|z|vec3|gravity, or omit for scalar) (term '" + spec + "')"); + } + + if (!type.empty()) { + term.has_binding = true; + if (type == "scalar" || type == "double") { + term.binding = Binding::kDouble; + } else if (type == "bool") { + term.binding = Binding::kBool; + } else if (type == "int") { + term.binding = Binding::kInt; + } else if (type == "size" || type == "size_t") { + term.binding = Binding::kSize; + } else if (type == "vector3" || type == "vec3") { + term.binding = Binding::kVector3; + } else if (type == "direction_vector") { + term.binding = Binding::kDirectionVector; + } else if (type == "quaternion") { + term.binding = Binding::kQuaternion; + } else { + throw std::invalid_argument( + "RlBridge: unknown type='" + type + "' (term '" + spec + "')"); + } + const bool scalar_binding = term.binding == Binding::kDouble || term.binding == Binding::kBool + || term.binding == Binding::kInt || term.binding == Binding::kSize; + const bool vector_binding = + term.binding == Binding::kVector3 || term.binding == Binding::kDirectionVector; + if (term.take == Take::kScalar && !scalar_binding) + throw std::invalid_argument( + "RlBridge: type=" + type + + " needs take=vec3 or take=x|y|z (a whole vector is not a scalar) (term '" + + spec + "')"); + if ((term.take == Take::kComponent || term.take == Take::kVector) && !vector_binding) + throw std::invalid_argument( + "RlBridge: take=" + take + + " needs a vector interface " + "(type=vector3|direction_vector), got type=" + + type + " (term '" + spec + "')"); + if (term.take == Take::kGravity && term.binding != Binding::kQuaternion) + throw std::invalid_argument( + "RlBridge: transform=projected_gravity needs " + "type=quaternion (term '" + + spec + "')"); + } + + validate_tokens( + tokens, {"path", "take", "transform", "type", "index", "scale", "clip", "default", "name"}, + spec); + return finish_common(term, tokens, spec); +} + +ActionTerm parse_action_term(const std::string& spec) { + ActionTerm term; + const auto tokens = parse_tokens(spec); + const auto index = tokens.find("index"); + if (index == tokens.end()) + throw std::invalid_argument("RlBridge: action term requires index= (term '" + spec + "')"); + const auto index_value = parse_integer(index->second, "action index"); + if (index_value < 0) + throw std::invalid_argument("RlBridge: action index must be >= 0 (term '" + spec + "')"); + term.index = static_cast(index_value); + + const auto output = tokens.find("output"); + if (output == tokens.end() || output->second.empty()) + throw std::invalid_argument( + "RlBridge: action term requires output= (term '" + spec + "')"); + term.output = output->second; + term.id = string_token(tokens, "name", term.output); + if (term.id.empty()) + throw std::invalid_argument("RlBridge: action name= must not be empty (term '" + spec + "')"); + term.scale = double_token(tokens, "scale", 1.0, spec); + if (term.scale == 0.0) + throw std::invalid_argument("RlBridge: action scale=0 is not allowed (term '" + spec + "')"); + parse_clip(tokens, spec, term.has_clip, term.clip_min, term.clip_max); + validate_tokens(tokens, {"index", "output", "name", "scale", "clip"}, spec); + return term; +} + +void assign_observation_indices(std::vector& terms) { + std::size_t cursor = 0; + for (auto& term : terms) { + if (term.has_index && term.index != cursor) + throw std::invalid_argument( + "RlBridge: observation term '" + term.id + + "' declares " + "index=" + + std::to_string(term.index) + " but the running offset is " + + std::to_string(cursor) + " (index= is an assertion, not a reorder)"); + term.index = cursor; + cursor += term.dim; + } +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/term_parser.hpp b/src/rl_bridge/term_parser.hpp new file mode 100644 index 0000000..042b56c --- /dev/null +++ b/src/rl_bridge/term_parser.hpp @@ -0,0 +1,24 @@ +#pragma once + +#include +#include +#include + +#include + +#include "rl_bridge/joint_config.hpp" +#include "rl_bridge/types.hpp" + +namespace rmcs_rl { + +struct TermParseContext { + rclcpp::Node& node; + const JointConfig& joints; + std::size_t action_size; +}; + +ObsTerm parse_obs_term(const std::string& spec, const TermParseContext& context); +ActionTerm parse_action_term(const std::string& spec); +void assign_observation_indices(std::vector& terms); + +} // namespace rmcs_rl diff --git a/src/rl_bridge/types.hpp b/src/rl_bridge/types.hpp new file mode 100644 index 0000000..705cbe9 --- /dev/null +++ b/src/rl_bridge/types.hpp @@ -0,0 +1,104 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace rmcs_rl { + +inline constexpr std::size_t kNoSlot = std::numeric_limits::max(); + +enum class TermKind { kPath, kJointPos, kJointVel, kJointTorque, kLastAction, kConstant }; + +enum class Take { kScalar, kComponent, kVector, kGravity }; + +enum class Binding { + kDouble, + kBool, + kInt, + kSize, + kVector3, + kDirectionVector, + kQuaternion, +}; + +enum class InvalidMode { kNaN, kZero, kHold }; + +struct ObsTerm { + TermKind kind = TermKind::kPath; + Take take = Take::kScalar; + + std::string path; + std::string id; + std::size_t dim = 1; + + std::size_t index = 0; + bool has_index = false; + int component = 0; + bool has_binding = false; + Binding binding = Binding::kDouble; + + std::vector joints; + std::vector joint_names; + std::vector joint_slots; + std::vector joint_defaults; + bool relative = false; + bool zero = false; + + std::vector action_indices; + std::vector constants; + + double scale = 1.0; + bool has_clip = false; + double clip_min = 0.0; + double clip_max = 0.0; + bool has_default = false; + double default_value = 0.0; + + std::size_t slot = kNoSlot; +}; + +struct ActionTerm { + std::size_t index = 0; + std::string output; + std::string id; + double scale = 1.0; + bool has_clip = false; + double clip_min = 0.0; + double clip_max = 0.0; +}; + +struct Slot { + std::string path; + Binding binding = Binding::kDouble; + bool required = true; + std::unique_ptr> double_value; + std::unique_ptr> bool_value; + std::unique_ptr> int_value; + std::unique_ptr> size_value; + std::unique_ptr> vector3_value; + std::unique_ptr< + rmcs_executor::Component::InputInterface> + direction_vector_value; + std::unique_ptr> quaternion_value; +}; + +struct ActionSnapshot { + std::uint64_t obs_seq = 0; + std::uint64_t layout_hash = 0; + std::uint64_t model_id = 0; + std::vector action; + std::chrono::steady_clock::time_point received{}; +}; + +} // namespace rmcs_rl diff --git a/src/rl_bridge/utility.cpp b/src/rl_bridge/utility.cpp new file mode 100644 index 0000000..fbb682a --- /dev/null +++ b/src/rl_bridge/utility.cpp @@ -0,0 +1,225 @@ +#include "rl_bridge/utility.hpp" + +#include +#include +#include +#include +#include + +#include +#include + +namespace rmcs_rl { + +std::vector split_by(const std::string& text, char delimiter) { + std::vector parts; + std::string current; + std::istringstream stream{text}; + while (std::getline(stream, current, delimiter)) + parts.push_back(current); + return parts; +} + +std::vector split_whitespace(const std::string& text) { + std::vector tokens; + std::istringstream stream{text}; + std::string token; + while (stream >> token) + tokens.push_back(token); + return tokens; +} + +std::string format_number(double value) { + char buffer[32]; + std::snprintf(buffer, sizeof(buffer), "%.6g", value); + return buffer; +} + +std::string join(const std::vector& parts, const char* separator) { + std::string text; + for (std::size_t i = 0; i < parts.size(); ++i) { + if (i != 0) + text += separator; + text += parts[i]; + } + return text; +} + +double parse_double(const std::string& text, const std::string& context) { + try { + std::size_t consumed = 0; + const double value = std::stod(text, &consumed); + if (consumed != text.size()) + throw std::invalid_argument("trailing characters"); + return value; + } catch (const std::exception&) { + throw std::invalid_argument("RlBridge: '" + text + "' is not a number (" + context + ")"); + } +} + +std::int64_t parse_integer(const std::string& text, const std::string& context) { + try { + std::size_t consumed = 0; + const auto value = std::stoll(text, &consumed); + if (consumed != text.size()) + throw std::invalid_argument("trailing characters"); + return value; + } catch (const std::exception&) { + throw std::invalid_argument("RlBridge: '" + text + "' is not an integer (" + context + ")"); + } +} + +bool parse_boolean(const std::string& text, const std::string& context) { + if (text == "true" || text == "1") + return true; + if (text == "false" || text == "0") + return false; + throw std::invalid_argument( + "RlBridge: '" + text + "' is not a boolean (true/false) (" + context + ")"); +} + +bool is_finite(double value) { return std::isfinite(value); } + +std::string pretty_type(const std::type_info& type) { + if (type == typeid(double)) + return "double"; + if (type == typeid(float)) + return "float"; + if (type == typeid(bool)) + return "bool"; + if (type == typeid(int)) + return "int"; + if (type == typeid(std::size_t)) + return "std::size_t"; + if (type == typeid(Eigen::Vector3d)) + return "Eigen::Vector3d"; + if (type == typeid(rmcs_description::BaseLink::DirectionVector)) + return "rmcs_description::BaseLink::DirectionVector"; + if (type == typeid(Eigen::Quaterniond)) + return "Eigen::Quaterniond"; + return std::string{type.name()} + " (unknown)"; +} + +std::map parse_tokens(const std::string& spec) { + std::map tokens; + for (const auto& token : split_whitespace(spec)) { + const auto equals = token.find('='); + if (equals == std::string::npos || equals == 0) + throw std::invalid_argument( + "RlBridge: term token '" + token + "' must be key=value (term '" + spec + "')"); + const std::string key = token.substr(0, equals); + const std::string value = token.substr(equals + 1); + if (!tokens.emplace(key, value).second) + throw std::invalid_argument( + "RlBridge: duplicate key '" + key + "' in term '" + spec + "'"); + } + return tokens; +} + +void validate_tokens( + const std::map& tokens, const std::vector& allowed, + const std::string& spec) { + for (const auto& [key, ignored] : tokens) { + (void)ignored; + if (std::find(allowed.begin(), allowed.end(), key) == allowed.end()) + throw std::invalid_argument( + "RlBridge: unknown key '" + key + "' in term '" + spec + "'"); + } +} + +double double_token( + const std::map& tokens, const std::string& key, double fallback, + const std::string& spec) { + const auto token = tokens.find(key); + if (token == tokens.end()) + return fallback; + const double value = parse_double(token->second, key + " in term '" + spec + "'"); + if (!is_finite(value)) + throw std::invalid_argument( + "RlBridge: " + key + " must be finite (term '" + spec + "')"); + return value; +} + +bool bool_token( + const std::map& tokens, const std::string& key, bool fallback, + const std::string& spec) { + const auto token = tokens.find(key); + if (token == tokens.end()) + return fallback; + return parse_boolean(token->second, key + " in term '" + spec + "'"); +} + +std::string string_token( + const std::map& tokens, const std::string& key, + const std::string& fallback) { + const auto token = tokens.find(key); + if (token == tokens.end()) + return fallback; + return token->second; +} + +void parse_clip( + const std::map& tokens, const std::string& spec, bool& has_clip, + double& clip_min, double& clip_max) { + const auto clip = tokens.find("clip"); + if (clip == tokens.end()) + return; + const auto separator = clip->second.find(':'); + if (separator == std::string::npos) { + const double symmetric = parse_double(clip->second, "clip"); + if (!(symmetric > 0.0) || !is_finite(symmetric)) + throw std::invalid_argument( + "RlBridge: single-value clip must be finite and > 0 (term '" + spec + "')"); + clip_min = -symmetric; + clip_max = symmetric; + } else { + clip_min = parse_double(clip->second.substr(0, separator), "clip min"); + clip_max = parse_double(clip->second.substr(separator + 1), "clip max"); + } + if (!is_finite(clip_min) || !is_finite(clip_max) || clip_min > clip_max) + throw std::invalid_argument("RlBridge: invalid clip range (term '" + spec + "')"); + has_clip = true; +} + +void parse_index( + const std::map& tokens, const std::string& spec, bool& has_index, + std::size_t& index) { + const auto token = tokens.find("index"); + if (token == tokens.end()) + return; + const auto value = parse_integer(token->second, "index in term '" + spec + "'"); + if (value < 0) + throw std::invalid_argument("RlBridge: index must be >= 0 (term '" + spec + "')"); + has_index = true; + index = static_cast(value); +} + +std::vector +parse_index_list(const std::string& text, std::size_t limit, const std::string& what) { + std::vector indices; + for (const auto& piece : split_by(text, ',')) { + if (piece.empty()) + throw std::invalid_argument("RlBridge: empty entry in " + what + " '" + text + "'"); + const auto range = piece.find(".."); + if (range == std::string::npos) { + const auto index = parse_integer(piece, what); + if (index < 0 || static_cast(index) >= limit) + throw std::invalid_argument( + "RlBridge: " + what + " index " + std::to_string(index) + + " out of range [0, " + std::to_string(limit) + ")"); + indices.push_back(static_cast(index)); + } else { + const auto first = parse_integer(piece.substr(0, range), what); + const auto last = parse_integer(piece.substr(range + 2), what); + if (first < 0 || last < first || static_cast(last) >= limit) + throw std::invalid_argument( + "RlBridge: invalid " + what + " range '" + piece + "' (limit " + + std::to_string(limit) + ")"); + for (auto index = first; index <= last; ++index) + indices.push_back(static_cast(index)); + } + } + return indices; +} + +} // namespace rmcs_rl diff --git a/src/rl_bridge/utility.hpp b/src/rl_bridge/utility.hpp new file mode 100644 index 0000000..96503fc --- /dev/null +++ b/src/rl_bridge/utility.hpp @@ -0,0 +1,44 @@ +#pragma once + +#include +#include +#include +#include +#include +#include + +namespace rmcs_rl { + +std::vector split_by(const std::string& text, char delimiter); +std::vector split_whitespace(const std::string& text); +std::string format_number(double value); +std::string join(const std::vector& parts, const char* separator); +double parse_double(const std::string& text, const std::string& context); +std::int64_t parse_integer(const std::string& text, const std::string& context); +bool parse_boolean(const std::string& text, const std::string& context); +bool is_finite(double value); +std::string pretty_type(const std::type_info& type); + +std::map parse_tokens(const std::string& spec); +void validate_tokens( + const std::map& tokens, const std::vector& allowed, + const std::string& spec); +double double_token( + const std::map& tokens, const std::string& key, double fallback, + const std::string& spec); +bool bool_token( + const std::map& tokens, const std::string& key, bool fallback, + const std::string& spec); +std::string string_token( + const std::map& tokens, const std::string& key, + const std::string& fallback); +void parse_clip( + const std::map& tokens, const std::string& spec, bool& has_clip, + double& clip_min, double& clip_max); +void parse_index( + const std::map& tokens, const std::string& spec, bool& has_index, + std::size_t& index); +std::vector +parse_index_list(const std::string& text, std::size_t limit, const std::string& what); + +} // namespace rmcs_rl diff --git a/tool/check_policy_contract.py b/tool/check_policy_contract.py index 3bc113a..1c8ef8c 100644 --- a/tool/check_policy_contract.py +++ b/tool/check_policy_contract.py @@ -198,7 +198,7 @@ def main() -> None: try: obs_terms, act_terms, obs_size, act_size = layout.load_config(args.config, args.node) obs_sig = layout.obs_signature( - obs_terms, act_size, layout.history_length(config, node)) + obs_terms, act_size, layout.history_length(args.config, args.node)) act_sig = layout.action_signature(act_terms) except layout.LayoutError as exc: print(f"FAIL config: {exc}", file=sys.stderr) From e7819853cc83a4d9175a6fe31749322e7d74e1c0 Mon Sep 17 00:00:00 2001 From: ZGZ713912 Date: Thu, 24 Sep 2026 21:41:11 +0800 Subject: [PATCH 3/3] refactor(rl_bridge): improve binding candidate selection --- src/rl_bridge.cpp | 43 ++++++++++++++++++----------- src/rl_bridge/interface_binding.hpp | 1 + 2 files changed, 28 insertions(+), 16 deletions(-) diff --git a/src/rl_bridge.cpp b/src/rl_bridge.cpp index 3c13d83..45d6a6a 100644 --- a/src/rl_bridge.cpp +++ b/src/rl_bridge.cpp @@ -283,15 +283,23 @@ class RlBridge final } case TermKind::kPath: { std::vector candidates; - switch (term.take) { - case Take::kScalar: - candidates = {Binding::kDouble, Binding::kBool, Binding::kInt, Binding::kSize}; - break; - case Take::kComponent: - case Take::kVector: - candidates = {Binding::kVector3, Binding::kDirectionVector}; - break; - case Take::kGravity: candidates = {Binding::kQuaternion}; break; + if (term.has_binding) { + candidates = {term.binding}; + } else { + switch (term.take) { + case Take::kScalar: + candidates = { + Binding::kDouble, + Binding::kBool, + Binding::kInt, + Binding::kSize}; + break; + case Take::kComponent: + case Take::kVector: + candidates = {Binding::kVector3, Binding::kDirectionVector}; + break; + case Take::kGravity: candidates = {Binding::kQuaternion}; break; + } } term.slot = acquire_slot( *this, term.path, candidates, !term.has_default, output_map, slots_, @@ -454,12 +462,20 @@ class RlBridge final prev_pub_seq_ = 0; } + // Interface storage is kept before configuration and runtime state so the component wiring + // remains visible at the class boundary. + std::vector slots_; + std::vector>> action_outputs_; + + OutputInterface valid_output_{}; + OutputInterface healthy_output_{}; + OutputInterface action_age_output_{}; + OutputInterface obs_seq_output_{}; + JointConfig joint_config_; std::vector obs_terms_; std::vector action_terms_; - std::vector slots_; - std::vector>> action_outputs_; std::unordered_set own_output_paths_; std::size_t obs_size_ = 0; @@ -501,11 +517,6 @@ class RlBridge final ActionChannel action_channel_; ActionSnapshot read_snapshot_; - OutputInterface valid_output_{}; - OutputInterface healthy_output_{}; - OutputInterface action_age_output_{}; - OutputInterface obs_seq_output_{}; - rclcpp::Publisher::SharedPtr obs_publisher_; rclcpp::Subscription::SharedPtr action_subscription_; }; diff --git a/src/rl_bridge/interface_binding.hpp b/src/rl_bridge/interface_binding.hpp index 9224374..23890b7 100644 --- a/src/rl_bridge/interface_binding.hpp +++ b/src/rl_bridge/interface_binding.hpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include