返回教程正文

配套源码

protocol.py

app/bricks/ros_gateway/protocol.py
ros_gateway Brick:在 App Lab 中建立可靠的 WebSocket 通道app/bricks/ros_gateway/protocol.py
Python247 行
  1. # SPDX-License-Identifier: MIT
  2. import json
  3. import math
  4. import time
  5. PROTOCOL_VERSION = 1
  6. ALLOWED_MODES = {"IDLE", "ROS_TELEOP", "ESTOP"}
  7. FUTURE_TOLERANCE_MS = 2000
  8. class ProtocolError(ValueError):
  9. """WebSocket 消息协议错误。"""
  10. def __init__(self, code, message, seq=None):
  11. """
  12. @description : 创建包含机器可读错误码的协议异常
  13. @param code : 机器可读错误码
  14. @param message : 面向日志的错误说明
  15. @param seq : 可选的原始消息序号
  16. @return : 无返回值
  17. """
  18. super().__init__(message)
  19. self.code = code
  20. self.seq = seq
  21. def now_ms():
  22. """
  23. @description : 获取 Unix 毫秒时间戳
  24. @param : 无参数
  25. @return : 当前 Unix 毫秒时间戳
  26. """
  27. return time.time_ns() // 1_000_000
  28. def decode_message(raw_message):
  29. """
  30. @description : 将文本 WebSocket 帧解析为基础协议消息
  31. @param raw_message : WebSocket 接收到的文本或 UTF-8 字节数据
  32. @return : 已完成版本和类型检查的字典
  33. """
  34. if isinstance(raw_message, bytes):
  35. try:
  36. raw_message = raw_message.decode("utf-8")
  37. except UnicodeDecodeError as exc:
  38. raise ProtocolError("invalid_encoding", "message must be UTF-8") from exc
  39. if not isinstance(raw_message, str):
  40. raise ProtocolError("invalid_frame", "message must be a text frame")
  41. try:
  42. message = json.loads(raw_message)
  43. except json.JSONDecodeError as exc:
  44. raise ProtocolError("invalid_json", f"invalid JSON: {exc.msg}") from exc
  45. if not isinstance(message, dict):
  46. raise ProtocolError("invalid_message", "JSON root must be an object")
  47. seq = message.get("seq") if type(message.get("seq")) is int else None
  48. if type(message.get("version")) is not int or message["version"] != PROTOCOL_VERSION:
  49. raise ProtocolError(
  50. "unsupported_version",
  51. f"version must be {PROTOCOL_VERSION}",
  52. seq,
  53. )
  54. message_type = message.get("type")
  55. if not isinstance(message_type, str) or not message_type:
  56. raise ProtocolError("invalid_type", "type must be a non-empty string", seq)
  57. return message
  58. def validate_hello(message):
  59. """
  60. @description : 校验 ROS 2 客户端握手消息
  61. @param message : 已解析的基础协议消息
  62. @return : 客户端节点名称
  63. """
  64. if message["type"] != "hello":
  65. raise ProtocolError("hello_required", "first message must be hello")
  66. if message.get("role") != "ros2":
  67. raise ProtocolError("invalid_role", "hello role must be ros2")
  68. node = message.get("node")
  69. if not isinstance(node, str) or not node.strip() or len(node) > 128:
  70. raise ProtocolError("invalid_node", "hello node must contain 1 to 128 characters")
  71. return node.strip()
  72. def validate_sequence(message, last_sequence):
  73. """
  74. @description : 校验命令序号为非负且严格递增
  75. @param message : 已解析的协议消息
  76. @param last_sequence : 当前连接最近接受的序号
  77. @return : 已校验的新序号
  78. """
  79. sequence = message.get("seq")
  80. if type(sequence) is not int or sequence < 0:
  81. raise ProtocolError("invalid_seq", "seq must be a non-negative integer")
  82. if sequence <= last_sequence:
  83. raise ProtocolError(
  84. "non_monotonic_seq",
  85. f"seq must be greater than {last_sequence}",
  86. sequence,
  87. )
  88. return sequence
  89. def validate_timestamp(message, maximum_age_ms, current_timestamp_ms=None):
  90. """
  91. @description : 校验消息时间戳并拒绝过期或明显来自未来的消息
  92. @param message : 已解析的协议消息
  93. @param maximum_age_ms : 允许的最大消息年龄,None 表示不检查过期
  94. @param current_timestamp_ms : 测试时可注入的当前 Unix 毫秒时间戳
  95. @return : 已校验的消息时间戳
  96. """
  97. timestamp = message.get("timestamp_ms")
  98. sequence = message.get("seq") if type(message.get("seq")) is int else None
  99. if type(timestamp) is not int or timestamp <= 0:
  100. raise ProtocolError(
  101. "invalid_timestamp",
  102. "timestamp_ms must be a positive integer",
  103. sequence,
  104. )
  105. current = now_ms() if current_timestamp_ms is None else current_timestamp_ms
  106. if timestamp > current + FUTURE_TOLERANCE_MS:
  107. raise ProtocolError(
  108. "future_timestamp",
  109. "timestamp_ms is too far in the future",
  110. sequence,
  111. )
  112. if maximum_age_ms is not None and current - timestamp > maximum_age_ms:
  113. raise ProtocolError(
  114. "stale_command",
  115. f"message age exceeds {maximum_age_ms} ms",
  116. sequence,
  117. )
  118. return timestamp
  119. def validate_heartbeat(message):
  120. """
  121. @description : 校验应用层心跳消息
  122. @param message : 已解析的协议消息
  123. @return : 心跳时间戳
  124. """
  125. if message["type"] != "heartbeat":
  126. raise ProtocolError("invalid_type", "message type must be heartbeat")
  127. return validate_timestamp(message, None)
  128. def validate_mode_change(message):
  129. """
  130. @description : 校验模式切换请求
  131. @param message : 已解析的协议消息
  132. @return : 已校验的目标模式
  133. """
  134. if message["type"] != "mode_change":
  135. raise ProtocolError("invalid_type", "message type must be mode_change")
  136. validate_timestamp(message, None)
  137. mode = message.get("mode")
  138. if mode not in ALLOWED_MODES:
  139. raise ProtocolError(
  140. "invalid_mode",
  141. f"mode must be one of {sorted(ALLOWED_MODES)}",
  142. message.get("seq"),
  143. )
  144. return mode
  145. def validate_cmd_vel(message, mode, limits, command_timeout_ms, current_timestamp_ms=None):
  146. """
  147. @description : 校验速度指令字段、模式、范围和时效性
  148. @param message : 已解析的 cmd_vel 消息
  149. @param mode : 当前底盘模式
  150. @param limits : vx、vy、wz 的绝对值上限
  151. @param command_timeout_ms : 允许的最大指令年龄
  152. @param current_timestamp_ms : 测试时可注入的当前 Unix 毫秒时间戳
  153. @return : 规范化后的速度指令字典
  154. """
  155. if message["type"] != "cmd_vel":
  156. raise ProtocolError("invalid_type", "message type must be cmd_vel")
  157. timestamp = validate_timestamp(message, command_timeout_ms, current_timestamp_ms)
  158. values = {}
  159. for field_name in ("vx", "vy", "wz"):
  160. field_value = message.get(field_name)
  161. if isinstance(field_value, bool) or not isinstance(field_value, (int, float)):
  162. raise ProtocolError(
  163. "invalid_field",
  164. f"{field_name} must be a finite number",
  165. message.get("seq"),
  166. )
  167. field_value = float(field_value)
  168. if not math.isfinite(field_value):
  169. raise ProtocolError(
  170. "invalid_field",
  171. f"{field_name} must be a finite number",
  172. message.get("seq"),
  173. )
  174. if abs(field_value) > limits[field_name]:
  175. raise ProtocolError(
  176. "out_of_range",
  177. f"abs({field_name}) must be <= {limits[field_name]}",
  178. message.get("seq"),
  179. )
  180. values[field_name] = field_value
  181. if mode != "ROS_TELEOP" and any(value != 0.0 for value in values.values()):
  182. raise ProtocolError(
  183. "mode_denied",
  184. "non-zero cmd_vel requires ROS_TELEOP mode",
  185. message.get("seq"),
  186. )
  187. return {
  188. "seq": message["seq"],
  189. "timestamp_ms": timestamp,
  190. **values,
  191. }
  192. def build_message(message_type, sequence=None, timestamp_ms=None, **fields):
  193. """
  194. @description : 构造统一版本的出站协议消息
  195. @param message_type : 消息类型
  196. @param sequence : 可选的消息序号
  197. @param timestamp_ms : 可选的 Unix 毫秒时间戳
  198. @param fields : 需要附加的消息字段
  199. @return : 可直接 JSON 序列化的消息字典
  200. """
  201. message = {
  202. "version": PROTOCOL_VERSION,
  203. "type": message_type,
  204. "timestamp_ms": now_ms() if timestamp_ms is None else timestamp_ms,
  205. }
  206. if sequence is not None:
  207. message["seq"] = sequence
  208. message.update(fields)
  209. return message