返回教程正文

配套源码

gateway.py

app/bricks/ros_gateway/gateway.py
ros_gateway Brick:在 App Lab 中建立可靠的 WebSocket 通道app/bricks/ros_gateway/gateway.py
Python599 行
  1. # SPDX-License-Identifier: MIT
  2. import json
  3. import os
  4. import queue
  5. import threading
  6. import time
  7. from arduino.app_utils import brick
  8. from websockets.exceptions import ConnectionClosed
  9. from websockets.sync.server import serve
  10. from .protocol import (
  11. ProtocolError,
  12. build_message,
  13. decode_message,
  14. validate_cmd_vel,
  15. validate_heartbeat,
  16. validate_hello,
  17. validate_mode_change,
  18. validate_sequence,
  19. )
  20. @brick
  21. class RosGateway:
  22. """App Lab 与原生 ROS 2 之间的 WebSocket 网关。"""
  23. def __init__(
  24. self,
  25. host=None,
  26. port=None,
  27. path=None,
  28. max_vx=None,
  29. max_vy=None,
  30. max_wz=None,
  31. command_timeout_ms=None,
  32. heartbeat_timeout_ms=None,
  33. outbound_queue_size=32,
  34. ):
  35. """
  36. @description : 创建 ROS Gateway 并读取 App Lab 注入的配置变量
  37. @param host : 容器内监听地址,None 时读取环境变量
  38. @param port : WebSocket 监听端口,None 时读取环境变量
  39. @param path : WebSocket 请求路径,None 时读取环境变量
  40. @param max_vx : vx 绝对值上限
  41. @param max_vy : vy 绝对值上限
  42. @param max_wz : wz 绝对值上限
  43. @param command_timeout_ms : 速度指令超时毫秒数
  44. @param heartbeat_timeout_ms : 连接心跳超时毫秒数
  45. @param outbound_queue_size : 有界出站队列长度
  46. @return : 无返回值
  47. """
  48. self._host = host or os.getenv("ROS_GATEWAY_HOST", "0.0.0.0")
  49. self._port = self._read_int(port, "ROS_GATEWAY_PORT", 8765, 1, 65535)
  50. self._path = path or os.getenv("ROS_GATEWAY_PATH", "/ros")
  51. self._limits = {
  52. "vx": self._read_float(max_vx, "ROS_GATEWAY_MAX_VX", 0.8, 0.0),
  53. "vy": self._read_float(max_vy, "ROS_GATEWAY_MAX_VY", 0.8, 0.0),
  54. "wz": self._read_float(max_wz, "ROS_GATEWAY_MAX_WZ", 1.5, 0.0),
  55. }
  56. self._command_timeout_ms = self._read_int(
  57. command_timeout_ms,
  58. "ROS_GATEWAY_COMMAND_TIMEOUT_MS",
  59. 300,
  60. 1,
  61. 60_000,
  62. )
  63. self._heartbeat_timeout_ms = self._read_int(
  64. heartbeat_timeout_ms,
  65. "ROS_GATEWAY_HEARTBEAT_TIMEOUT_MS",
  66. 3000,
  67. self._command_timeout_ms,
  68. 120_000,
  69. )
  70. self._outbound = queue.Queue(maxsize=max(1, int(outbound_queue_size)))
  71. self._state_lock = threading.RLock()
  72. self._stop_event = threading.Event()
  73. self._ready_event = threading.Event()
  74. self._server = None
  75. self._server_thread = None
  76. self._active_connection = None
  77. self._connected = False
  78. self._client_node = None
  79. self._mode = "IDLE"
  80. self._last_sequence = -1
  81. self._outbound_sequence = 0
  82. self._last_rx_monotonic = 0.0
  83. self._last_cmd_monotonic = 0.0
  84. self._last_server_heartbeat = 0.0
  85. self._watchdog_triggered = True
  86. self._server_error = None
  87. self._dropped_messages = 0
  88. self._cmd_vel_callback = None
  89. self._mode_change_callback = None
  90. self._stop_callback = None
  91. def start(self):
  92. """
  93. @description : 启动非阻塞 WebSocket 服务线程
  94. @param : 无参数
  95. @return : 无返回值
  96. """
  97. if self._server_thread and self._server_thread.is_alive():
  98. return
  99. self._stop_event.clear()
  100. self._ready_event.clear()
  101. self._server_thread = threading.Thread(
  102. target=self._run_server,
  103. name="ros-gateway-server",
  104. daemon=True,
  105. )
  106. self._server_thread.start()
  107. if not self._ready_event.wait(timeout=5.0):
  108. print("[ros_gateway] server did not become ready within 5 seconds", flush=True)
  109. def stop(self):
  110. """
  111. @description : 关闭连接、释放端口并等待服务线程退出
  112. @param : 无参数
  113. @return : 无返回值
  114. """
  115. self._stop_event.set()
  116. with self._state_lock:
  117. server = self._server
  118. if server is not None:
  119. server.shutdown(reason="App Lab application is stopping")
  120. if self._server_thread and self._server_thread is not threading.current_thread():
  121. self._server_thread.join(timeout=5.0)
  122. self._invoke_stop("app_shutdown")
  123. def on_cmd_vel(self, callback):
  124. """
  125. @description : 注册已通过安全校验的 cmd_vel 回调
  126. @param callback : 接收规范化速度字典的函数
  127. @return : 当前 RosGateway 实例
  128. """
  129. self._cmd_vel_callback = callback
  130. return self
  131. def on_mode_change(self, callback):
  132. """
  133. @description : 注册模式切换回调
  134. @param callback : 接收目标模式并返回是否允许的函数
  135. @return : 当前 RosGateway 实例
  136. """
  137. self._mode_change_callback = callback
  138. return self
  139. def on_stop(self, callback):
  140. """
  141. @description : 注册通信异常或命令超时的统一安全停车回调
  142. @param callback : 接收停车原因字符串的函数
  143. @return : 当前 RosGateway 实例
  144. """
  145. self._stop_callback = callback
  146. return self
  147. def is_ros_connected(self):
  148. """
  149. @description : 查询是否存在已完成握手的 ROS 2 客户端
  150. @param : 无参数
  151. @return : 已连接返回 True,否则返回 False
  152. """
  153. with self._state_lock:
  154. return self._connected
  155. def get_status(self):
  156. """
  157. @description : 获取连接、模式、队列和服务错误状态快照
  158. @param : 无参数
  159. @return : 状态字典
  160. """
  161. with self._state_lock:
  162. return {
  163. "connected": self._connected,
  164. "client_node": self._client_node,
  165. "mode": self._mode,
  166. "server_ready": self._ready_event.is_set(),
  167. "server_error": self._server_error,
  168. "queued_messages": self._outbound.qsize(),
  169. "dropped_messages": self._dropped_messages,
  170. }
  171. def publish_base_state(self, state):
  172. """
  173. @description : 校验并将模拟或真实底盘状态加入有界发送队列
  174. @param state : 底盘状态字段字典
  175. @return : 成功入队返回 True,无客户端或失败返回 False
  176. """
  177. required_fields = {
  178. "mode",
  179. "enabled",
  180. "wheel_position",
  181. "wheel_velocity",
  182. "battery_voltage",
  183. "estop",
  184. "fault_code",
  185. }
  186. missing_fields = required_fields - state.keys()
  187. if missing_fields:
  188. raise ValueError(f"base_state missing fields: {sorted(missing_fields)}")
  189. self._validate_four_values("wheel_position", state["wheel_position"])
  190. self._validate_four_values("wheel_velocity", state["wheel_velocity"])
  191. return self._enqueue("base_state", **state)
  192. def publish_imu(self, imu):
  193. """
  194. @description : 将 IMU 状态加入有界发送队列
  195. @param imu : orientation、angular_velocity、linear_acceleration 字典
  196. @return : 成功入队返回 True,无客户端或失败返回 False
  197. """
  198. self._validate_vector("orientation", imu.get("orientation"), 4)
  199. self._validate_vector("angular_velocity", imu.get("angular_velocity"), 3)
  200. self._validate_vector("linear_acceleration", imu.get("linear_acceleration"), 3)
  201. return self._enqueue("imu", **imu)
  202. def publish_diagnostics(self, diagnostics):
  203. """
  204. @description : 将诊断字典加入有界发送队列
  205. @param diagnostics : 可 JSON 序列化的诊断字段
  206. @return : 成功入队返回 True,无客户端或失败返回 False
  207. """
  208. if not isinstance(diagnostics, dict):
  209. raise ValueError("diagnostics must be a dictionary")
  210. return self._enqueue("diagnostics", **diagnostics)
  211. @staticmethod
  212. def _read_int(explicit_value, environment_name, default_value, minimum, maximum):
  213. """
  214. @description : 从显式参数或环境变量读取受范围约束的整数
  215. @param explicit_value : 调用方显式提供的值
  216. @param environment_name : 环境变量名称
  217. @param default_value : 默认值
  218. @param minimum : 最小允许值
  219. @param maximum : 最大允许值
  220. @return : 校验后的整数
  221. """
  222. value = explicit_value
  223. if value is None:
  224. value = os.getenv(environment_name, str(default_value))
  225. parsed = int(value)
  226. if parsed < minimum or parsed > maximum:
  227. raise ValueError(f"{environment_name} must be between {minimum} and {maximum}")
  228. return parsed
  229. @staticmethod
  230. def _read_float(explicit_value, environment_name, default_value, minimum):
  231. """
  232. @description : 从显式参数或环境变量读取有下限约束的浮点数
  233. @param explicit_value : 调用方显式提供的值
  234. @param environment_name : 环境变量名称
  235. @param default_value : 默认值
  236. @param minimum : 最小允许值
  237. @return : 校验后的浮点数
  238. """
  239. value = explicit_value
  240. if value is None:
  241. value = os.getenv(environment_name, str(default_value))
  242. parsed = float(value)
  243. if parsed <= minimum:
  244. raise ValueError(f"{environment_name} must be greater than {minimum}")
  245. return parsed
  246. @staticmethod
  247. def _validate_vector(field_name, value, expected_length):
  248. """
  249. @description : 校验固定长度数值数组
  250. @param field_name : 用于错误信息的字段名
  251. @param value : 待校验数组
  252. @param expected_length : 期望数组长度
  253. @return : 无返回值
  254. """
  255. if not isinstance(value, (list, tuple)) or len(value) != expected_length:
  256. raise ValueError(f"{field_name} must contain {expected_length} values")
  257. for item in value:
  258. if isinstance(item, bool) or not isinstance(item, (int, float)):
  259. raise ValueError(f"{field_name} must contain only numbers")
  260. @classmethod
  261. def _validate_four_values(cls, field_name, value):
  262. """
  263. @description : 校验四轮位置或速度数组
  264. @param field_name : 用于错误信息的字段名
  265. @param value : 待校验四元素数组
  266. @return : 无返回值
  267. """
  268. cls._validate_vector(field_name, value, 4)
  269. def _run_server(self):
  270. """
  271. @description : 在专用线程中绑定端口并运行同步 WebSocket 服务
  272. @param : 无参数
  273. @return : 无返回值
  274. """
  275. try:
  276. with serve(
  277. self._handle_connection,
  278. self._host,
  279. self._port,
  280. compression=None,
  281. ping_interval=1.0,
  282. ping_timeout=1.0,
  283. close_timeout=2.0,
  284. max_size=16 * 1024,
  285. max_queue=16,
  286. ) as server:
  287. with self._state_lock:
  288. self._server = server
  289. self._server_error = None
  290. self._ready_event.set()
  291. print(
  292. f"[ros_gateway] listening on ws://{self._host}:{self._port}{self._path}",
  293. flush=True,
  294. )
  295. server.serve_forever()
  296. except Exception as exc:
  297. with self._state_lock:
  298. self._server_error = f"{type(exc).__name__}: {exc}"
  299. self._ready_event.set()
  300. print(f"[ros_gateway] server failed: {type(exc).__name__}: {exc}", flush=True)
  301. finally:
  302. with self._state_lock:
  303. self._server = None
  304. self._ready_event.clear()
  305. def _handle_connection(self, websocket):
  306. """
  307. @description : 处理单个 WebSocket 客户端的握手、收发和安全清理
  308. @param websocket : websockets 同步服务端连接对象
  309. @return : 无返回值
  310. """
  311. request_path = websocket.request.path.split("?", 1)[0]
  312. if request_path != self._path:
  313. websocket.close(1008, f"expected path {self._path}")
  314. return
  315. with self._state_lock:
  316. if self._active_connection is not None:
  317. websocket.close(1013, "another ROS 2 client already owns the gateway")
  318. return
  319. self._active_connection = websocket
  320. was_connected = False
  321. disconnect_reason = "connection_closed"
  322. try:
  323. raw_hello = websocket.recv(timeout=2.0)
  324. hello = decode_message(raw_hello)
  325. client_node = validate_hello(hello)
  326. current = time.monotonic()
  327. with self._state_lock:
  328. self._connected = True
  329. self._client_node = client_node
  330. self._last_sequence = -1
  331. self._last_rx_monotonic = current
  332. self._last_cmd_monotonic = 0.0
  333. self._last_server_heartbeat = current
  334. self._watchdog_triggered = True
  335. was_connected = True
  336. self._send_json(
  337. websocket,
  338. build_message(
  339. "hello",
  340. role="app",
  341. node="ros-gateway-loopback",
  342. ),
  343. )
  344. print(f"[ros_gateway] ROS 2 connected: node={client_node}", flush=True)
  345. while not self._stop_event.is_set():
  346. self._service_timers(websocket)
  347. self._drain_outbound(websocket)
  348. try:
  349. raw_message = websocket.recv(timeout=0.05)
  350. except TimeoutError:
  351. continue
  352. self._process_message(websocket, raw_message)
  353. except ProtocolError as exc:
  354. disconnect_reason = exc.code
  355. self._send_error(websocket, exc)
  356. websocket.close(1008, str(exc))
  357. except ConnectionClosed as exc:
  358. disconnect_reason = f"connection_closed:{exc.code}"
  359. except Exception as exc:
  360. disconnect_reason = f"handler_error:{type(exc).__name__}"
  361. print(f"[ros_gateway] connection handler error: {type(exc).__name__}: {exc}", flush=True)
  362. finally:
  363. with self._state_lock:
  364. if self._active_connection is websocket:
  365. self._active_connection = None
  366. self._connected = False
  367. self._client_node = None
  368. self._mode = "IDLE"
  369. self._clear_outbound_queue()
  370. if was_connected:
  371. print(f"[ros_gateway] ROS 2 disconnected: {disconnect_reason}", flush=True)
  372. self._invoke_stop(disconnect_reason)
  373. def _process_message(self, websocket, raw_message):
  374. """
  375. @description : 分派一个客户端消息并执行对应协议校验
  376. @param websocket : 当前活动 WebSocket 连接
  377. @param raw_message : 原始文本消息
  378. @return : 无返回值
  379. """
  380. try:
  381. message = decode_message(raw_message)
  382. message_type = message["type"]
  383. if message_type == "hello":
  384. raise ProtocolError("duplicate_hello", "hello is only valid as the first message")
  385. with self._state_lock:
  386. sequence = validate_sequence(message, self._last_sequence)
  387. if message_type == "heartbeat":
  388. validate_heartbeat(message)
  389. with self._state_lock:
  390. self._last_sequence = sequence
  391. self._last_rx_monotonic = time.monotonic()
  392. return
  393. if message_type == "mode_change":
  394. target_mode = validate_mode_change(message)
  395. accepted = True
  396. if self._mode_change_callback is not None:
  397. accepted = self._mode_change_callback(target_mode) is not False
  398. if accepted:
  399. with self._state_lock:
  400. self._mode = target_mode
  401. with self._state_lock:
  402. self._last_sequence = sequence
  403. self._last_rx_monotonic = time.monotonic()
  404. self._send_json(
  405. websocket,
  406. build_message(
  407. "ack",
  408. sequence=sequence,
  409. command="mode_change",
  410. accepted=accepted,
  411. mode=self._mode,
  412. ),
  413. )
  414. return
  415. if message_type == "cmd_vel":
  416. with self._state_lock:
  417. mode = self._mode
  418. command = validate_cmd_vel(
  419. message,
  420. mode,
  421. self._limits,
  422. self._command_timeout_ms,
  423. )
  424. with self._state_lock:
  425. current = time.monotonic()
  426. self._last_sequence = sequence
  427. self._last_rx_monotonic = current
  428. self._last_cmd_monotonic = current
  429. self._watchdog_triggered = False
  430. if self._cmd_vel_callback is not None:
  431. self._cmd_vel_callback(command)
  432. return
  433. raise ProtocolError(
  434. "unknown_type",
  435. f"unsupported message type: {message_type}",
  436. sequence,
  437. )
  438. except ProtocolError as exc:
  439. self._send_error(websocket, exc)
  440. def _service_timers(self, websocket):
  441. """
  442. @description : 发送服务端心跳并执行连接与速度命令看门狗
  443. @param websocket : 当前活动 WebSocket 连接
  444. @return : 无返回值
  445. """
  446. current = time.monotonic()
  447. with self._state_lock:
  448. last_rx = self._last_rx_monotonic
  449. last_cmd = self._last_cmd_monotonic
  450. watchdog_triggered = self._watchdog_triggered
  451. last_heartbeat = self._last_server_heartbeat
  452. if current - last_rx > self._heartbeat_timeout_ms / 1000.0:
  453. raise ProtocolError("heartbeat_timeout", "client heartbeat timed out")
  454. if last_cmd > 0.0 and not watchdog_triggered:
  455. if current - last_cmd > self._command_timeout_ms / 1000.0:
  456. with self._state_lock:
  457. self._watchdog_triggered = True
  458. self._invoke_stop("cmd_vel_timeout")
  459. if current - last_heartbeat >= 1.0:
  460. self._send_json(websocket, self._next_message("heartbeat"))
  461. with self._state_lock:
  462. self._last_server_heartbeat = current
  463. def _enqueue(self, message_type, **fields):
  464. """
  465. @description : 将出站消息加入有界队列,满时丢弃最旧消息
  466. @param message_type : 出站消息类型
  467. @param fields : 出站消息字段
  468. @return : 有活动客户端时返回 True,否则返回 False
  469. """
  470. if not self.is_ros_connected():
  471. return False
  472. message = self._next_message(message_type, **fields)
  473. try:
  474. self._outbound.put_nowait(message)
  475. except queue.Full:
  476. try:
  477. self._outbound.get_nowait()
  478. except queue.Empty:
  479. pass
  480. self._outbound.put_nowait(message)
  481. with self._state_lock:
  482. self._dropped_messages += 1
  483. return True
  484. def _next_message(self, message_type, **fields):
  485. """
  486. @description : 分配严格递增的服务端序号并构造消息
  487. @param message_type : 出站消息类型
  488. @param fields : 出站消息字段
  489. @return : 已构造的消息字典
  490. """
  491. with self._state_lock:
  492. self._outbound_sequence += 1
  493. sequence = self._outbound_sequence
  494. return build_message(message_type, sequence=sequence, **fields)
  495. def _drain_outbound(self, websocket):
  496. """
  497. @description : 在连接处理线程中发送当前队列内的全部消息
  498. @param websocket : 当前活动 WebSocket 连接
  499. @return : 无返回值
  500. """
  501. while True:
  502. try:
  503. message = self._outbound.get_nowait()
  504. except queue.Empty:
  505. return
  506. self._send_json(websocket, message)
  507. def _send_error(self, websocket, error):
  508. """
  509. @description : 向客户端发送结构化协议错误
  510. @param websocket : 当前活动 WebSocket 连接
  511. @param error : ProtocolError 实例
  512. @return : 无返回值
  513. """
  514. payload = build_message(
  515. "error",
  516. code=error.code,
  517. message=str(error),
  518. )
  519. if error.seq is not None:
  520. payload["seq"] = error.seq
  521. self._send_json(websocket, payload)
  522. print(f"[ros_gateway] rejected message: {error.code}: {error}", flush=True)
  523. @staticmethod
  524. def _send_json(websocket, message):
  525. """
  526. @description : 将字典编码为紧凑 UTF-8 JSON 文本并发送
  527. @param websocket : 当前活动 WebSocket 连接
  528. @param message : 待发送消息字典
  529. @return : 无返回值
  530. """
  531. websocket.send(json.dumps(message, ensure_ascii=False, separators=(",", ":")))
  532. def _invoke_stop(self, reason):
  533. """
  534. @description : 安全调用统一停车回调并隔离回调异常
  535. @param reason : 停车原因
  536. @return : 无返回值
  537. """
  538. if self._stop_callback is None:
  539. return
  540. try:
  541. self._stop_callback(reason)
  542. except Exception as exc:
  543. print(f"[ros_gateway] stop callback failed: {type(exc).__name__}: {exc}", flush=True)
  544. def _clear_outbound_queue(self):
  545. """
  546. @description : 清空断线客户端尚未发送的出站消息
  547. @param : 无参数
  548. @return : 无返回值
  549. """
  550. while True:
  551. try:
  552. self._outbound.get_nowait()
  553. except queue.Empty:
  554. return