feat(lekiwi): release V0.10.1 初步集成 LeKiwi,优化碰撞模型
集成通用机器人数值接口、本机控制桥、LeRobot 插件和统一键盘遥操作。采用离线 CoACD 全臂碰撞配方 revision 4、局部装配区切分与结构自接触,限制直接关节位姿写入并保留安全看门狗。同步版本号、变更记录、来源许可证和兼容性验证。
This commit is contained in:
@@ -0,0 +1,202 @@
|
||||
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -0,0 +1,64 @@
|
||||
# 本机机器人控制桥接
|
||||
|
||||
独立于 `training_server` 的实时数值接口。默认监听 `127.0.0.1:8766`;不导入 LeRobot、PyTorch 或训练任务,不执行代码、Shell 或服务器文件路径。
|
||||
|
||||
## 启动
|
||||
|
||||
从仓库根目录执行(Python 3.12):
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
python -m pip install -e ./control_bridge
|
||||
python -m mujoco_control_bridge
|
||||
# 或 npm run control-bridge
|
||||
```
|
||||
|
||||
未设置 `MUJOCO_CONTROL_TOKEN` 时,服务生成至少 256 bit 随机 token,只在启动终端显示一次。也可通过该环境变量提供自己的 token(16–4096 个可打印 ASCII 字符,不含空白)。`--port 0` 自动分配端口,仅测试需要;不能配置非回环监听地址。
|
||||
|
||||
1. 在本机 HTTP 工作台加载受支持的机器人,显式选择 profile 并通过编译校验。
|
||||
2. 在“控制台 → 开源项目 / 外部控制”填写地址和 token,连接桥接。
|
||||
3. 播放,然后点击“允许外部控制”。接管会停止 Python/ONNX、锁定 1×并禁止冲突写入。
|
||||
4. 在另一终端给 SDK 设置相同 token(不要把 token 放进 URL、源码或命令历史):
|
||||
|
||||
```bash
|
||||
read -rsp '控制 token: ' MUJOCO_CONTROL_TOKEN; echo
|
||||
export MUJOCO_CONTROL_TOKEN
|
||||
export MUJOCO_CONTROL_ENDPOINT=http://127.0.0.1:8766
|
||||
```
|
||||
|
||||
## 通用 Python SDK
|
||||
|
||||
客户端路径仅使用 Python 标准库;HTTP 不使用环境代理,也不跟随重定向。接口不依赖 LeKiwi 的通道名称或数量。
|
||||
|
||||
```python
|
||||
import os
|
||||
import time
|
||||
from mujoco_control_bridge import SimRobotClient
|
||||
|
||||
with SimRobotClient(os.environ["MUJOCO_CONTROL_ENDPOINT"]) as robot:
|
||||
descriptor = robot.describe()
|
||||
values = {
|
||||
channel["id"]: min(channel["max"], max(channel["min"], 0.0))
|
||||
for channel in descriptor["actionChannels"]
|
||||
}
|
||||
start = time.monotonic()
|
||||
for i in range(90):
|
||||
measured = robot.get_observation() # 实际状态,不是目标回显
|
||||
accepted = robot.send_action(values) # 物理步后确认,包含已接受的限幅目标
|
||||
time.sleep(max(0, start + (i + 1) / 30 - time.monotonic()))
|
||||
```
|
||||
|
||||
通用 SDK 要求提交完整通道集合;部分动作的保持/默认语义由具体插件处理。`robot.reset()` 返回新 `modelEpoch` 的暂停观测,并清除租约;必须在浏览器重新播放、授权,再连接。`is_connected` 不是永久健康承诺:距成功动作超过 500 ms 会在本地失效。
|
||||
|
||||
## 边界与故障处理
|
||||
|
||||
- HTTP Bearer token;WS 首帧认证。校验本机 Host/Origin,拒绝查询字符串 token、远程页面和未经认证的预检/连接。
|
||||
- 一个浏览器后端、一个 Python 写入者;读取状态不取得写权限。不能自动抢占已有控制者。
|
||||
- 64 KiB 帧/请求、每租约至多 100 动作/秒、8 个在途 RPC、一个最新待应用动作;旧目标被替换时返回 `SUPERSEDED`,不会堆积无限队列。
|
||||
- 浏览器注册/认证有绝对超时;1 秒 RPC 确认期限、500 ms 动作/观测新鲜度边界。同序号心跳不刷新观测年龄。
|
||||
- 暂停、断连、进程崩溃、重载、页面隐藏、冻结超时均撤销授权。轮目标归零、臂/夹爪保持实测姿态,并暂停外控仿真;**不是把所有位置伺服置零**。
|
||||
- 回到页面或重新连接不会恢复旧命令。先检查错误、恢复关节限位,再播放/授权。
|
||||
- 无 HTTPS/WSS、跨机器访问、锁步、图像流、LeRobot 数据集或训练 API。不要使用端口转发或反向代理扩大此 V1 的信任边界。
|
||||
- 浏览器控制 token 仅在页面内存;不写入 `localStorage`/`sessionStorage`、工程导出或应用日志。拥有 token 的本机进程仍被视为可信控制者。
|
||||
|
||||
测试:`python -m unittest discover -s control_bridge/tests -v`。完整消息结构、错误码和扩展点见 [机器人接口](../docs/robot-interface.md);LeRobot 用法见 [LeKiwi 示例](../examples/lekiwi/README.md)。
|
||||
@@ -0,0 +1,18 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=77,<82"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "mujoco-control-bridge"
|
||||
version = "0.1.0"
|
||||
description = "Loopback-only numeric robot control bridge for MuJoCo Web"
|
||||
requires-python = ">=3.12"
|
||||
license = "Apache-2.0"
|
||||
license-files = ["LICENSE"]
|
||||
dependencies = ["aiohttp>=3.12,<4"]
|
||||
|
||||
[project.scripts]
|
||||
mujoco-control-bridge = "mujoco_control_bridge.__main__:main"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -0,0 +1,4 @@
|
||||
from .client import SimRobotClient
|
||||
from .protocol import RobotError
|
||||
|
||||
__all__ = ["RobotError", "SimRobotClient"]
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Run an independent bridge on loopback only; token never goes in URLs."""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import signal
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from .server import create_app
|
||||
|
||||
|
||||
async def serve(port, token):
|
||||
runner = web.AppRunner(create_app(token), access_log=None, shutdown_timeout=2)
|
||||
await runner.setup()
|
||||
try:
|
||||
await web.TCPSite(runner, "127.0.0.1", port).start()
|
||||
address = runner.addresses[0]
|
||||
print(
|
||||
"CONTROL_BRIDGE_READY " + json.dumps({"endpoint": f"http://127.0.0.1:{address[1]}"}),
|
||||
flush=True,
|
||||
)
|
||||
stopped = asyncio.Event()
|
||||
loop = asyncio.get_running_loop()
|
||||
for signum in (signal.SIGINT, signal.SIGTERM):
|
||||
loop.add_signal_handler(signum, stopped.set)
|
||||
await stopped.wait()
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="MuJoCo 本机数值控制桥接(不是训练服务)")
|
||||
parser.add_argument("--port", type=int, default=8766)
|
||||
args = parser.parse_args()
|
||||
if not 0 <= args.port <= 65535:
|
||||
parser.error("端口必须为 0–65535")
|
||||
token = os.environ.get("MUJOCO_CONTROL_TOKEN")
|
||||
if not token:
|
||||
token = secrets.token_urlsafe(32)
|
||||
print(f"本次控制 token(仅显示一次;复制到浏览器/Python):{token}", flush=True)
|
||||
asyncio.run(serve(args.port, token))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,160 @@
|
||||
"""Synchronous, generic robot SDK; stdlib only on the client path."""
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.parse import urlparse
|
||||
from urllib.request import HTTPRedirectHandler, ProxyHandler, Request, build_opener
|
||||
|
||||
from . import protocol as p
|
||||
from .protocol import RobotError
|
||||
|
||||
|
||||
class NoRedirect(HTTPRedirectHandler):
|
||||
def redirect_request(self, req, fp, code, msg, headers, newurl):
|
||||
return None # Never forward the bearer token to a redirected endpoint.
|
||||
|
||||
|
||||
class SimRobotClient:
|
||||
def __init__(self, endpoint="http://127.0.0.1:8766", token=None, timeout=2.0):
|
||||
parsed = urlparse(endpoint)
|
||||
if (
|
||||
parsed.scheme != "http"
|
||||
or parsed.hostname not in {"127.0.0.1", "localhost"}
|
||||
or parsed.username
|
||||
or parsed.password
|
||||
or parsed.query
|
||||
or parsed.fragment
|
||||
or parsed.path not in {"", "/"}
|
||||
):
|
||||
raise ValueError("V1 仅支持本机 http://127.0.0.1:port 或 localhost")
|
||||
self.endpoint = endpoint.rstrip("/")
|
||||
self.token = token or os.environ.get("MUJOCO_CONTROL_TOKEN", "")
|
||||
if not self.token:
|
||||
raise ValueError("需要 token 或 MUJOCO_CONTROL_TOKEN")
|
||||
if (
|
||||
not isinstance(self.token, str)
|
||||
or not 16 <= len(self.token) <= 4096
|
||||
or not all(33 <= ord(c) <= 126 for c in self.token)
|
||||
):
|
||||
raise ValueError("控制 token 必须为16–4096字符的可打印 ASCII,且不含空白")
|
||||
self.timeout = timeout
|
||||
self._opener = build_opener(ProxyHandler({}), NoRedirect())
|
||||
self._descriptor = None
|
||||
self._identity = None
|
||||
self._action_seq = 0
|
||||
self._last_action_at = 0.0
|
||||
|
||||
def _request(self, method, path, payload=None, lease=None):
|
||||
headers = {"Authorization": f"Bearer {self.token}"}
|
||||
if lease:
|
||||
headers["X-Control-Lease"] = lease["leaseId"]
|
||||
raw = None
|
||||
if payload is not None:
|
||||
raw = json.dumps(payload, allow_nan=False).encode()
|
||||
if len(raw) > 65536:
|
||||
raise RobotError("INVALID_MESSAGE", "消息超过64KiB")
|
||||
headers["Content-Type"] = "application/json"
|
||||
request = Request(
|
||||
f"{self.endpoint}/api/control/v1{path}", data=raw, headers=headers, method=method
|
||||
)
|
||||
try:
|
||||
with self._opener.open(request, timeout=self.timeout) as response:
|
||||
data = response.read(65537)
|
||||
if len(data) > 65536:
|
||||
raise RobotError("INVALID_MESSAGE", "响应超过64KiB")
|
||||
return p.loads(data)
|
||||
except HTTPError as exc:
|
||||
try:
|
||||
error = p.loads(exc.read(65537))["error"]
|
||||
code, message = error["code"], error["message"]
|
||||
except (RobotError, KeyError, TypeError):
|
||||
code, message = "DISCONNECTED", f"桥接 HTTP {exc.code}"
|
||||
raise RobotError(code, message) from None
|
||||
except (URLError, TimeoutError, ConnectionError, OSError) as exc:
|
||||
raise RobotError("DISCONNECTED", "无法访问本机控制桥接") from exc
|
||||
|
||||
@property
|
||||
def is_connected(self):
|
||||
return self._identity is not None and time.monotonic() - self._last_action_at <= 0.5
|
||||
|
||||
def describe(self):
|
||||
self._descriptor = p.descriptor(self._request("GET", "/robot"))
|
||||
return self._descriptor
|
||||
|
||||
def connect(self):
|
||||
if self._identity:
|
||||
raise RobotError("CONFLICT", "客户端已连接;请先 disconnect")
|
||||
desc = self.describe()
|
||||
obs = p.observation(self._request("GET", "/observation"), desc)
|
||||
claim = {
|
||||
"sessionId": obs["sessionId"],
|
||||
"modelEpoch": obs["modelEpoch"],
|
||||
"modelFingerprint": desc["modelFingerprint"],
|
||||
}
|
||||
lease = p.record(self._request("POST", "/lease", claim), p.IDENTITY)
|
||||
p.identity(lease)
|
||||
p.same_identity(lease, obs, ("sessionId", "modelEpoch"))
|
||||
self._identity = lease
|
||||
self._action_seq = 0
|
||||
self._last_action_at = time.monotonic()
|
||||
return desc
|
||||
|
||||
def _connected(self):
|
||||
if not self.is_connected:
|
||||
raise RobotError("DISCONNECTED", "没有有效租约;请重新在浏览器授权并连接")
|
||||
return dict(self._identity)
|
||||
|
||||
def get_observation(self):
|
||||
identity = self._connected()
|
||||
obs = p.observation(self._request("GET", "/observation"), self._descriptor)
|
||||
p.same_identity(obs, identity, ("sessionId", "modelEpoch"))
|
||||
return obs
|
||||
|
||||
def send_action(self, values):
|
||||
identity = self._connected()
|
||||
values = p.values(values, self._descriptor["actionChannels"])
|
||||
self._action_seq += 1
|
||||
packet = {"protocolVersion": 1, **identity, "actionSeq": self._action_seq, "values": values}
|
||||
try:
|
||||
result = p.action_result(
|
||||
self._request("POST", "/action", packet, lease=identity), self._descriptor
|
||||
)
|
||||
p.same_identity(result, packet, (*p.IDENTITY, "actionSeq"))
|
||||
self._last_action_at = time.monotonic()
|
||||
return result
|
||||
except RobotError:
|
||||
self.disconnect()
|
||||
raise
|
||||
|
||||
def reset(self):
|
||||
identity = self._connected()
|
||||
try:
|
||||
observed = p.observation(
|
||||
self._request("POST", "/reset", {}, lease=identity), self._descriptor
|
||||
)
|
||||
if (
|
||||
observed["sessionId"] != identity["sessionId"]
|
||||
or observed["modelEpoch"] <= identity["modelEpoch"]
|
||||
or not observed["paused"]
|
||||
):
|
||||
raise RobotError("STALE", "reset 没有返回新代次")
|
||||
return observed
|
||||
finally:
|
||||
self.disconnect()
|
||||
|
||||
def disconnect(self):
|
||||
identity, self._identity = self._identity, None
|
||||
if identity:
|
||||
# The server/browser watchdog remains the final safety barrier.
|
||||
with contextlib.suppress(RobotError):
|
||||
self._request("DELETE", "/lease", lease=identity)
|
||||
|
||||
def __enter__(self):
|
||||
self.connect()
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
self.disconnect()
|
||||
@@ -0,0 +1,196 @@
|
||||
"""Strict, framework-independent v1 wire validation (mirrors TS validation.ts)."""
|
||||
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
|
||||
UNITS = {"rad", "rad/s", "m", "m/s", "ratio", "N", "N.m"}
|
||||
MODES = {"position", "velocity", "effort", "opening"}
|
||||
IDENTITY = {"sessionId", "modelEpoch", "leaseId"}
|
||||
|
||||
|
||||
class RobotError(RuntimeError):
|
||||
def __init__(self, code, message):
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
|
||||
def as_dict(self):
|
||||
return {"code": self.code, "message": str(self)}
|
||||
|
||||
|
||||
def invalid(message):
|
||||
raise RobotError("INVALID_MESSAGE", message)
|
||||
|
||||
|
||||
def record(value, keys, label="message"):
|
||||
if not isinstance(value, dict) or set(value) != set(keys):
|
||||
invalid(f"{label} 字段不完整或包含未知字段")
|
||||
return value
|
||||
|
||||
|
||||
def text(value):
|
||||
if not isinstance(value, str) or not re.fullmatch(r"[a-zA-Z0-9][a-zA-Z0-9_.:-]{0,127}", value):
|
||||
invalid("标识符无效")
|
||||
return value
|
||||
|
||||
|
||||
def finite(value):
|
||||
try:
|
||||
good = type(value) in (int, float) and math.isfinite(value)
|
||||
except OverflowError:
|
||||
good = False
|
||||
if not good:
|
||||
invalid("数值必须有限且不能为 bool")
|
||||
return value
|
||||
|
||||
|
||||
def integer(value, minimum=0):
|
||||
finite(value)
|
||||
if int(value) != value or not minimum <= value <= 2**53 - 1:
|
||||
invalid("整数超出安全范围")
|
||||
return value
|
||||
|
||||
|
||||
def version(value):
|
||||
if type(value) not in (int, float) or value != 1:
|
||||
raise RobotError("UNSUPPORTED", "只支持协议 v1")
|
||||
|
||||
|
||||
def identity(value):
|
||||
text(value["sessionId"])
|
||||
integer(value["modelEpoch"])
|
||||
text(value["leaseId"])
|
||||
|
||||
|
||||
def channels(value):
|
||||
if not isinstance(value, list) or not 0 < len(value) <= 256:
|
||||
invalid("通道列表无效")
|
||||
seen = set()
|
||||
for channel in value:
|
||||
record(channel, {"id", "unit", "mode", "min", "max"}, "channel")
|
||||
name = text(channel["id"])
|
||||
if name in seen:
|
||||
invalid("通道名称无效或重复")
|
||||
seen.add(name)
|
||||
if str(channel["unit"]) not in UNITS or str(channel["mode"]) not in MODES:
|
||||
invalid("未知单位或通道模式")
|
||||
for key in ("min", "max"):
|
||||
finite(channel[key])
|
||||
if channel["min"] >= channel["max"]:
|
||||
invalid("通道范围无效")
|
||||
return value
|
||||
|
||||
|
||||
def descriptor(value):
|
||||
record(
|
||||
value,
|
||||
{
|
||||
"protocolVersion",
|
||||
"profileId",
|
||||
"profileVersion",
|
||||
"modelFingerprint",
|
||||
"frame",
|
||||
"actionChannels",
|
||||
"observationChannels",
|
||||
"capabilities",
|
||||
},
|
||||
"descriptor",
|
||||
)
|
||||
version(value["protocolVersion"])
|
||||
if value["frame"] != "x-forward-y-left-z-up":
|
||||
raise RobotError("UNSUPPORTED", "坐标系不支持")
|
||||
text(value["profileId"])
|
||||
integer(value["profileVersion"], 1)
|
||||
if not isinstance(value["modelFingerprint"], str) or not re.fullmatch(
|
||||
r"[a-f0-9]{64}", value["modelFingerprint"]
|
||||
):
|
||||
invalid("需要 SHA-256 模型指纹")
|
||||
channels(value["actionChannels"])
|
||||
channels(value["observationChannels"])
|
||||
caps = record(
|
||||
value["capabilities"], {"reset", "lockstep", "cameras", "training"}, "capabilities"
|
||||
)
|
||||
if any(type(v) is not bool for v in caps.values()) or any(
|
||||
caps[k] is not False for k in ("lockstep", "cameras", "training")
|
||||
):
|
||||
raise RobotError("UNSUPPORTED", "V1 不支持相机、锁步或训练")
|
||||
return value
|
||||
|
||||
|
||||
def values(value, specs, clamp=False):
|
||||
record(value, {c["id"] for c in specs}, "values")
|
||||
result = {}
|
||||
for channel in specs:
|
||||
number = finite(value[channel["id"]])
|
||||
result[channel["id"]] = (
|
||||
max(channel["min"], min(channel["max"], number)) if clamp else number
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def action(value, desc):
|
||||
record(value, {"protocolVersion", *IDENTITY, "actionSeq", "values"}, "action")
|
||||
version(value["protocolVersion"])
|
||||
identity(value)
|
||||
integer(value["actionSeq"], 1)
|
||||
return {**value, "values": values(value["values"], desc["actionChannels"], True)}
|
||||
|
||||
|
||||
def action_result(value, desc):
|
||||
record(value, {*IDENTITY, "actionSeq", "values", "simTime"}, "action result")
|
||||
identity(value)
|
||||
integer(value["actionSeq"], 1)
|
||||
finite(value["simTime"])
|
||||
if value["simTime"] < 0:
|
||||
invalid("仿真时间必须非负")
|
||||
values(value["values"], desc["actionChannels"])
|
||||
return value
|
||||
|
||||
|
||||
def observation(value, desc):
|
||||
record(
|
||||
value,
|
||||
{
|
||||
"protocolVersion",
|
||||
"sessionId",
|
||||
"modelEpoch",
|
||||
"sequence",
|
||||
"simTime",
|
||||
"appliedActionSeq",
|
||||
"paused",
|
||||
"values",
|
||||
},
|
||||
"observation",
|
||||
)
|
||||
version(value["protocolVersion"])
|
||||
text(value["sessionId"])
|
||||
integer(value["modelEpoch"])
|
||||
integer(value["sequence"])
|
||||
integer(value["appliedActionSeq"])
|
||||
finite(value["simTime"])
|
||||
if value["simTime"] < 0 or type(value["paused"]) is not bool:
|
||||
invalid("仿真时钟/暂停标记无效")
|
||||
values(value["values"], desc["observationChannels"])
|
||||
return value
|
||||
|
||||
|
||||
def same_identity(actual, expected, keys=IDENTITY):
|
||||
if any(actual.get(key) != expected.get(key) for key in keys):
|
||||
raise RobotError("STALE", "会话、模型代次或控制租约已失效")
|
||||
|
||||
|
||||
def loads(raw):
|
||||
def pairs(items):
|
||||
result = {}
|
||||
for key, value in items:
|
||||
if key in result:
|
||||
invalid("重复 JSON 字段")
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
try:
|
||||
return json.loads(
|
||||
raw, parse_constant=lambda _: invalid("禁止 NaN/Infinity"), object_pairs_hook=pairs
|
||||
)
|
||||
except (ValueError, TypeError, RecursionError) as exc:
|
||||
raise RobotError("INVALID_MESSAGE", "无效 JSON") from exc
|
||||
@@ -0,0 +1,463 @@
|
||||
"""Independent loopback broker. No simulation, training or LeRobot imports."""
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import hmac
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
from aiohttp import WSMsgType, web
|
||||
|
||||
from . import protocol as p
|
||||
from .protocol import RobotError
|
||||
|
||||
MAX_BYTES = 65536
|
||||
PREFIX = "/api/control/v1"
|
||||
ORIGIN = re.compile(r"http://(?:127\.0\.0\.1|localhost)(?::[0-9]{1,5})?\Z")
|
||||
HOST = re.compile(r"(?:127\.0\.0\.1|localhost)(?::[0-9]{1,5})?\Z")
|
||||
BROKER = web.AppKey("broker", object)
|
||||
|
||||
|
||||
class Broker:
|
||||
def __init__(self, token):
|
||||
if (
|
||||
not isinstance(token, str)
|
||||
or not 16 <= len(token) <= 4096
|
||||
or not all(33 <= ord(c) <= 126 for c in token)
|
||||
):
|
||||
raise ValueError("控制 token 必须为16–4096字符的可打印 ASCII,且不含空白")
|
||||
self.token = token
|
||||
self.ws = None
|
||||
self.descriptor = None
|
||||
self.observation = None
|
||||
self.observed_at = 0.0
|
||||
self.authorization_generation = 0
|
||||
self.blocked_generation = -1
|
||||
self.authorized = False
|
||||
self.lease = None
|
||||
self.last_action_at = 0.0
|
||||
self.last_action_seq = 0
|
||||
self.rate = deque()
|
||||
self.pending = {}
|
||||
|
||||
def robot(self):
|
||||
if self.ws is None or self.ws.closed or self.descriptor is None:
|
||||
raise RobotError("DISCONNECTED", "浏览器机器人未连接")
|
||||
return self.descriptor
|
||||
|
||||
def fresh_observation(self):
|
||||
self.robot()
|
||||
# Paused samples stop advancing by design; do not hide the required play
|
||||
# step behind a freshness error merely because the user waited to connect.
|
||||
if self.observation is not None and self.observation["paused"]:
|
||||
raise RobotError("PAUSED", "仿真已暂停;请先在浏览器点击“播放”,再点击“允许外部控制”")
|
||||
if self.observation is None or time.monotonic() - self.observed_at > 0.5:
|
||||
raise RobotError(
|
||||
"STALE",
|
||||
"机器人观测超过 500ms 未更新;请保持仿真页面可见,"
|
||||
"检查仿真是否卡顿,再播放并重新允许外部控制",
|
||||
)
|
||||
return self.observation
|
||||
|
||||
async def state(self, packet):
|
||||
p.record(packet, {"type", "observation", "enabled", "authorizationGeneration"})
|
||||
if type(packet["enabled"]) is not bool:
|
||||
p.invalid("授权标记必须为 bool")
|
||||
generation = p.integer(packet["authorizationGeneration"])
|
||||
if generation < self.authorization_generation:
|
||||
raise RobotError("STALE", "授权代次倒退")
|
||||
observed = p.observation(packet["observation"], self.robot())
|
||||
old = self.observation
|
||||
if old:
|
||||
p.same_identity(observed, old, ("sessionId",))
|
||||
if observed["modelEpoch"] < old["modelEpoch"] or (
|
||||
observed["modelEpoch"] == old["modelEpoch"]
|
||||
and observed["sequence"] < old["sequence"]
|
||||
):
|
||||
raise RobotError("STALE", "观测代次/序号倒退")
|
||||
# Sequence may restart only in a new epoch. Same-sample heartbeats aren't fresh.
|
||||
if (
|
||||
old is None
|
||||
or observed["modelEpoch"] > old["modelEpoch"]
|
||||
or observed["sequence"] > old["sequence"]
|
||||
):
|
||||
self.observation = observed
|
||||
self.observed_at = time.monotonic()
|
||||
if self.lease and generation != self.authorization_generation:
|
||||
# Cancel the previous owner without consuming the new explicit grant.
|
||||
await self.revoke("浏览器已重新授权", notify=False)
|
||||
self.authorization_generation = generation
|
||||
self.authorized = packet["enabled"] and generation > self.blocked_generation
|
||||
if self.lease and (
|
||||
not self.authorized
|
||||
or observed["paused"]
|
||||
or any(observed[k] != self.lease[k] for k in ("sessionId", "modelEpoch"))
|
||||
):
|
||||
await self.revoke("浏览器已撤销控制/重置模型", notify=False)
|
||||
|
||||
async def rpc(self, operation, payload):
|
||||
self.robot()
|
||||
if operation == "action":
|
||||
for key, (future, op) in list(self.pending.items()):
|
||||
if op == "action":
|
||||
if not future.done():
|
||||
future.set_exception(RobotError("SUPERSEDED", "已由更新的目标替代"))
|
||||
self.pending.pop(key, None)
|
||||
if len(self.pending) >= 8:
|
||||
raise RobotError("CONFLICT", "待确认请求已满")
|
||||
request_id = secrets.token_hex(16)
|
||||
future = asyncio.get_running_loop().create_future()
|
||||
ws, lease, generation = self.ws, self.lease, self.authorization_generation
|
||||
self.pending[request_id] = (future, operation)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self.ws.send_json(
|
||||
{"type": "request", "id": request_id, "op": operation, "payload": payload}
|
||||
),
|
||||
1,
|
||||
)
|
||||
return await asyncio.wait_for(future, 1)
|
||||
except TimeoutError as exc:
|
||||
if (
|
||||
self.ws is ws
|
||||
and self.lease == lease
|
||||
and self.authorization_generation == generation
|
||||
):
|
||||
await self.revoke("动作应用确认超时")
|
||||
raise RobotError("TIMEOUT", "浏览器未在一秒内确认应用") from exc
|
||||
finally:
|
||||
self.pending.pop(request_id, None)
|
||||
if not future.done():
|
||||
future.cancel()
|
||||
elif not future.cancelled():
|
||||
future.exception() # Also consume failures delivered during a blocked send.
|
||||
|
||||
async def revoke(self, reason, notify=True):
|
||||
self.lease = None
|
||||
self.authorized = False
|
||||
self.blocked_generation = self.authorization_generation
|
||||
for future, _ in self.pending.values():
|
||||
if not future.done():
|
||||
future.set_exception(RobotError("DISCONNECTED", reason))
|
||||
self.pending.clear()
|
||||
if notify and self.ws is not None and not self.ws.closed and self.observation:
|
||||
with contextlib.suppress(ConnectionError, TimeoutError):
|
||||
await asyncio.wait_for(
|
||||
self.ws.send_json(
|
||||
{
|
||||
"type": "stop",
|
||||
"reason": reason,
|
||||
"sessionId": self.observation["sessionId"],
|
||||
"authorizationGeneration": self.authorization_generation,
|
||||
}
|
||||
),
|
||||
0.1,
|
||||
)
|
||||
|
||||
def require_lease(self, request):
|
||||
provided = request.headers.get("X-Control-Lease", "")
|
||||
if (
|
||||
not self.lease
|
||||
or not provided.isascii()
|
||||
or not hmac.compare_digest(provided.encode(), self.lease["leaseId"].encode())
|
||||
):
|
||||
raise RobotError("UNAUTHORIZED", "控制租约无效")
|
||||
return dict(self.lease)
|
||||
|
||||
async def monitor(self):
|
||||
while True:
|
||||
await asyncio.sleep(0.05)
|
||||
if self.lease and (
|
||||
time.monotonic() - self.last_action_at > 0.5
|
||||
or time.monotonic() - self.observed_at > 0.5
|
||||
):
|
||||
await self.revoke("动作或观测看门狗超时 (500ms)")
|
||||
|
||||
|
||||
@web.middleware
|
||||
async def security(request, handler):
|
||||
broker = request.app[BROKER]
|
||||
try:
|
||||
if not HOST.fullmatch(request.headers.get("Host", "")):
|
||||
raise RobotError("UNAUTHORIZED", "Host 不被允许")
|
||||
origin = request.headers.get("Origin")
|
||||
if origin is not None and not ORIGIN.fullmatch(origin):
|
||||
raise RobotError("UNAUTHORIZED", "Origin 不被允许")
|
||||
if request.query:
|
||||
raise RobotError("INVALID_MESSAGE", "禁止 URL 查询参数及 URL 中的 token")
|
||||
if request.path != "/ws/control/v1":
|
||||
expected = f"Bearer {broker.token}".encode()
|
||||
provided = request.headers.get("Authorization", "")
|
||||
if not provided.isascii() or not hmac.compare_digest(provided.encode(), expected):
|
||||
raise RobotError("UNAUTHORIZED", "需要本机控制 Bearer token")
|
||||
elif origin is None:
|
||||
raise RobotError("UNAUTHORIZED", "浏览器 WebSocket 必须提供 Origin")
|
||||
return await handler(request)
|
||||
except RobotError as exc:
|
||||
status = {
|
||||
"UNAUTHORIZED": 401,
|
||||
"DISCONNECTED": 503,
|
||||
"TIMEOUT": 504,
|
||||
"INVALID_MESSAGE": 400,
|
||||
"UNSUPPORTED": 400,
|
||||
}.get(exc.code, 409)
|
||||
return web.json_response({"error": exc.as_dict()}, status=status)
|
||||
except web.HTTPRequestEntityTooLarge:
|
||||
return web.json_response(
|
||||
{"error": {"code": "INVALID_MESSAGE", "message": "消息不能超过64KiB"}}, status=413
|
||||
)
|
||||
|
||||
|
||||
async def body(request):
|
||||
if request.content_type != "application/json":
|
||||
p.invalid("需要 application/json")
|
||||
return p.loads(await request.read())
|
||||
|
||||
|
||||
async def health(request):
|
||||
b = request.app[BROKER]
|
||||
return web.json_response(
|
||||
{
|
||||
"protocolVersion": 1,
|
||||
"backendConnected": b.descriptor is not None,
|
||||
"authorized": b.authorized,
|
||||
"hasLease": b.lease is not None,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def robot(request):
|
||||
return web.json_response(request.app[BROKER].robot())
|
||||
|
||||
|
||||
async def observation(request):
|
||||
return web.json_response(request.app[BROKER].fresh_observation())
|
||||
|
||||
|
||||
async def claim(request):
|
||||
b = request.app[BROKER]
|
||||
data = p.record(await body(request), {"sessionId", "modelEpoch", "modelFingerprint"})
|
||||
p.text(data["sessionId"])
|
||||
p.integer(data["modelEpoch"])
|
||||
b.robot()
|
||||
observed = b.fresh_observation()
|
||||
p.same_identity(data, observed, ("sessionId", "modelEpoch"))
|
||||
if data["modelFingerprint"] != b.descriptor["modelFingerprint"]:
|
||||
raise RobotError("INCOMPATIBLE_MODEL", "模型指纹已变化")
|
||||
if b.lease:
|
||||
raise RobotError("CONFLICT", "已有一个 Python 控制者")
|
||||
if not b.authorized:
|
||||
raise RobotError("UNAUTHORIZED", "请先在浏览器显式允许外部控制")
|
||||
lease = {
|
||||
"sessionId": data["sessionId"],
|
||||
"modelEpoch": data["modelEpoch"],
|
||||
"leaseId": secrets.token_hex(24),
|
||||
}
|
||||
b.lease = lease
|
||||
b.last_action_at = time.monotonic()
|
||||
b.last_action_seq = 0
|
||||
b.rate.clear()
|
||||
try:
|
||||
result = p.record(
|
||||
await b.rpc(
|
||||
"claim",
|
||||
{
|
||||
**lease,
|
||||
"modelFingerprint": data["modelFingerprint"],
|
||||
"authorizationGeneration": b.authorization_generation,
|
||||
},
|
||||
),
|
||||
p.IDENTITY,
|
||||
)
|
||||
p.same_identity(result, lease)
|
||||
if b.lease != lease:
|
||||
raise RobotError("STALE", "租约在申请期间失效")
|
||||
return web.json_response(result)
|
||||
except BaseException:
|
||||
if b.lease == lease:
|
||||
await b.revoke("租约申请失败")
|
||||
raise
|
||||
|
||||
|
||||
async def action(request):
|
||||
b = request.app[BROKER]
|
||||
lease = b.require_lease(request)
|
||||
b.fresh_observation()
|
||||
data = p.action(await body(request), b.robot())
|
||||
p.same_identity(data, lease)
|
||||
if data["actionSeq"] <= b.last_action_seq:
|
||||
raise RobotError("STALE", "拒绝重复/乱序动作")
|
||||
now = time.monotonic()
|
||||
while b.rate and now - b.rate[0] >= 1:
|
||||
b.rate.popleft()
|
||||
if len(b.rate) >= 100:
|
||||
raise RobotError("CONFLICT", "动作频率不能超过100Hz")
|
||||
b.rate.append(now)
|
||||
b.last_action_seq = data["actionSeq"]
|
||||
b.last_action_at = now
|
||||
try:
|
||||
result = p.action_result(await b.rpc("action", data), b.robot())
|
||||
p.same_identity(result, data, (*p.IDENTITY, "actionSeq"))
|
||||
except RobotError as exc:
|
||||
if exc.code != "SUPERSEDED" and b.lease == lease:
|
||||
await b.revoke("动作确认失败")
|
||||
raise
|
||||
if b.lease != lease:
|
||||
raise RobotError("STALE", "动作确认来自失效租约")
|
||||
return web.json_response(result)
|
||||
|
||||
|
||||
async def release(request):
|
||||
b = request.app[BROKER]
|
||||
lease = b.require_lease(request)
|
||||
try:
|
||||
await b.rpc("release", lease)
|
||||
finally:
|
||||
if b.lease == lease:
|
||||
await b.revoke("控制者断开连接")
|
||||
return web.json_response({"released": True})
|
||||
|
||||
|
||||
async def reset(request):
|
||||
b = request.app[BROKER]
|
||||
lease = b.require_lease(request)
|
||||
p.record(await body(request), set())
|
||||
if not b.robot()["capabilities"]["reset"]:
|
||||
raise RobotError("UNSUPPORTED", "机器人不支持 reset")
|
||||
try:
|
||||
result = p.observation(await b.rpc("reset", lease), b.robot())
|
||||
if (
|
||||
result["sessionId"] != lease["sessionId"]
|
||||
or result["modelEpoch"] <= lease["modelEpoch"]
|
||||
or not result["paused"]
|
||||
):
|
||||
p.invalid("reset 未返回新代次的暂停观测")
|
||||
finally:
|
||||
if b.lease == lease:
|
||||
await b.revoke("reset 后需要重新授权")
|
||||
return web.json_response(result)
|
||||
|
||||
|
||||
async def websocket(request):
|
||||
b = request.app[BROKER]
|
||||
ws = web.WebSocketResponse(
|
||||
max_msg_size=MAX_BYTES, heartbeat=2, receive_timeout=5, compress=False
|
||||
)
|
||||
await ws.prepare(request)
|
||||
registered = False
|
||||
times = deque()
|
||||
try:
|
||||
# Outer deadline: WS ping/pong must NOT restart the authentication clock.
|
||||
first = await asyncio.wait_for(ws.receive(), 5)
|
||||
if first.type != WSMsgType.TEXT:
|
||||
raise RobotError("UNAUTHORIZED", "必须在5秒内通过认证")
|
||||
auth = p.record(p.loads(first.data), {"type", "token", "protocolVersion"})
|
||||
p.version(auth["protocolVersion"])
|
||||
if (
|
||||
auth["type"] != "auth"
|
||||
or not isinstance(auth["token"], str)
|
||||
or not auth["token"].isascii()
|
||||
or not hmac.compare_digest(auth["token"].encode(), b.token.encode())
|
||||
):
|
||||
raise RobotError("UNAUTHORIZED", "WebSocket token 无效")
|
||||
if b.ws is not None:
|
||||
raise RobotError("CONFLICT", "已有浏览器后端连接")
|
||||
b.ws = ws # Reserve the single backend before any further await.
|
||||
await ws.send_json({"type": "authenticated"})
|
||||
while True:
|
||||
message = (
|
||||
await asyncio.wait_for(ws.receive(), 5) if not registered else await ws.receive()
|
||||
)
|
||||
if message.type != WSMsgType.TEXT:
|
||||
break
|
||||
now = time.monotonic()
|
||||
while times and now - times[0] > 1:
|
||||
times.popleft()
|
||||
if len(times) >= 256:
|
||||
raise RobotError("CONFLICT", "浏览器消息过于频繁")
|
||||
times.append(now)
|
||||
packet = p.loads(message.data)
|
||||
if not isinstance(packet, dict):
|
||||
p.invalid("消息必须是对象")
|
||||
kind = packet.get("type")
|
||||
if kind == "register" and not registered:
|
||||
p.record(
|
||||
packet,
|
||||
{"type", "descriptor", "observation", "enabled", "authorizationGeneration"},
|
||||
)
|
||||
b.descriptor = p.descriptor(packet["descriptor"])
|
||||
await b.state({k: v for k, v in packet.items() if k != "descriptor"})
|
||||
registered = True
|
||||
await ws.send_json({"type": "ready"})
|
||||
elif kind == "state" and registered:
|
||||
await b.state(packet)
|
||||
elif kind == "result" and registered:
|
||||
p.record(packet, {"type", "id", "ok", "value"})
|
||||
p.text(packet["id"])
|
||||
if type(packet["ok"]) is not bool:
|
||||
p.invalid("ok 必须为 bool")
|
||||
pending = b.pending.get(packet["id"])
|
||||
if pending and not pending[0].done():
|
||||
if packet["ok"]:
|
||||
pending[0].set_result(packet["value"])
|
||||
else:
|
||||
error = p.record(packet["value"], {"code", "message"})
|
||||
if (
|
||||
not isinstance(error["code"], str)
|
||||
or not isinstance(error["message"], str)
|
||||
or len(error["message"]) > 1024
|
||||
):
|
||||
p.invalid("错误格式无效")
|
||||
pending[0].set_exception(RobotError(error["code"], error["message"]))
|
||||
else:
|
||||
p.invalid("未注册后端或未知消息类型")
|
||||
except (RobotError, TimeoutError, ConnectionError) as exc:
|
||||
error = (
|
||||
exc
|
||||
if isinstance(exc, RobotError)
|
||||
else RobotError("DISCONNECTED", "浏览器连接中断/超时")
|
||||
)
|
||||
if not ws.closed:
|
||||
with contextlib.suppress(ConnectionError):
|
||||
await ws.send_json({"type": "error", "error": error.as_dict()})
|
||||
finally:
|
||||
if b.ws is ws:
|
||||
await b.revoke("浏览器已断开", notify=False)
|
||||
b.ws = b.descriptor = b.observation = None
|
||||
b.authorization_generation = 0
|
||||
b.blocked_generation = -1
|
||||
await ws.close()
|
||||
return ws
|
||||
|
||||
|
||||
def create_app(token):
|
||||
app = web.Application(middlewares=[security], client_max_size=MAX_BYTES)
|
||||
app[BROKER] = Broker(token)
|
||||
app.add_routes(
|
||||
[
|
||||
web.get(f"{PREFIX}/health", health),
|
||||
web.get(f"{PREFIX}/robot", robot),
|
||||
web.get(f"{PREFIX}/observation", observation),
|
||||
web.post(f"{PREFIX}/lease", claim),
|
||||
web.delete(f"{PREFIX}/lease", release),
|
||||
web.post(f"{PREFIX}/action", action),
|
||||
web.post(f"{PREFIX}/reset", reset),
|
||||
web.get("/ws/control/v1", websocket),
|
||||
]
|
||||
)
|
||||
|
||||
async def lifetime(application):
|
||||
b = application[BROKER]
|
||||
task = asyncio.create_task(b.monitor())
|
||||
yield
|
||||
task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await task
|
||||
await b.revoke("桥接服务关闭")
|
||||
if b.ws is not None:
|
||||
await b.ws.close()
|
||||
|
||||
app.cleanup_ctx.append(lifetime)
|
||||
return app
|
||||
@@ -0,0 +1,65 @@
|
||||
import copy
|
||||
import json
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from mujoco_control_bridge import protocol as p
|
||||
from mujoco_control_bridge.client import SimRobotClient
|
||||
|
||||
FIXTURE = json.loads(
|
||||
(Path(__file__).resolve().parents[2] / "contracts/fixtures/single-joint.json").read_text()
|
||||
)
|
||||
|
||||
|
||||
class ProtocolTests(unittest.TestCase):
|
||||
def test_shared_fixture(self):
|
||||
descriptor = p.descriptor(FIXTURE["descriptor"])
|
||||
p.observation(FIXTURE["observation"], descriptor)
|
||||
action = p.action(FIXTURE["validAction"], descriptor)
|
||||
p.same_identity(action, FIXTURE["identity"])
|
||||
for invalid in FIXTURE["invalidActions"]:
|
||||
with self.subTest(invalid["label"]), self.assertRaises(p.RobotError) as caught:
|
||||
action = p.action({**FIXTURE["validAction"], **invalid["patch"]}, descriptor)
|
||||
p.same_identity(action, FIXTURE["identity"])
|
||||
self.assertEqual(caught.exception.code, invalid["code"])
|
||||
|
||||
def test_finite_clamp_and_extra_fields(self):
|
||||
for value in (float("nan"), float("inf"), True, "0", 10**1000):
|
||||
with self.subTest(value=str(value)[:20]), self.assertRaises(p.RobotError):
|
||||
p.values({"slider.position": value}, FIXTURE["descriptor"]["actionChannels"])
|
||||
self.assertEqual(
|
||||
p.values({"slider.position": 20}, FIXTURE["descriptor"]["actionChannels"], True),
|
||||
{"slider.position": 1},
|
||||
)
|
||||
for raw in ('{"x":1,"x":2}', '{"x":NaN}', '{"x":Infinity}', "["):
|
||||
with self.assertRaises(p.RobotError):
|
||||
p.loads(raw)
|
||||
|
||||
def test_unknown_units_capabilities_versions(self):
|
||||
for mutation in ("unit", "frame", "capability", "fingerprint"):
|
||||
descriptor = copy.deepcopy(FIXTURE["descriptor"])
|
||||
if mutation == "unit":
|
||||
descriptor["actionChannels"][0]["unit"] = "deg"
|
||||
elif mutation == "frame":
|
||||
descriptor["frame"] = "unknown"
|
||||
elif mutation == "fingerprint":
|
||||
descriptor["modelFingerprint"] = "not-sha"
|
||||
else:
|
||||
descriptor["capabilities"]["training"] = True
|
||||
with self.subTest(mutation), self.assertRaises(p.RobotError):
|
||||
p.descriptor(descriptor)
|
||||
|
||||
def test_only_local_endpoints_and_token_not_in_url(self):
|
||||
for endpoint in (
|
||||
"https://127.0.0.1",
|
||||
"http://evil.test",
|
||||
"http://user:token@127.0.0.1",
|
||||
"http://localhost/?token=secret",
|
||||
"http://localhost/path",
|
||||
):
|
||||
with self.subTest(endpoint), self.assertRaises(ValueError):
|
||||
SimRobotClient(endpoint, token="test-only-not-a-real-secret")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,322 @@
|
||||
import asyncio
|
||||
import contextlib
|
||||
import copy
|
||||
import json
|
||||
import time
|
||||
import unittest
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
|
||||
from aiohttp import ClientSession, WSMsgType, WSServerHandshakeError
|
||||
from aiohttp.test_utils import TestServer
|
||||
from mujoco_control_bridge import RobotError, SimRobotClient
|
||||
from mujoco_control_bridge.server import BROKER, PREFIX, create_app
|
||||
|
||||
FIXTURE = json.loads(
|
||||
(Path(__file__).resolve().parents[2] / "contracts/fixtures/single-joint.json").read_text()
|
||||
)
|
||||
TOKEN = "unit-test-token-not-for-real-use"
|
||||
ORIGIN = "http://127.0.0.1:5173"
|
||||
|
||||
|
||||
class Backend:
|
||||
"""One-channel numeric peer: no LeKiwi knowledge, never used for physics acceptance."""
|
||||
|
||||
def __init__(self, session, url):
|
||||
self.session, self.url = session, url
|
||||
self.obs = copy.deepcopy(FIXTURE["observation"])
|
||||
self.enabled = True
|
||||
self.generation = 1
|
||||
self.acknowledge = True
|
||||
self.calls = []
|
||||
|
||||
async def connect(self):
|
||||
self.ws = await self.session.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN)
|
||||
await self.ws.send_json({"type": "auth", "token": TOKEN, "protocolVersion": 1})
|
||||
assert (await self.ws.receive_json())["type"] == "authenticated"
|
||||
await self.ws.send_json(
|
||||
{**self.state(), "type": "register", "descriptor": FIXTURE["descriptor"]}
|
||||
)
|
||||
assert (await self.ws.receive_json())["type"] == "ready"
|
||||
self.task = asyncio.create_task(self.consume())
|
||||
|
||||
def state(self):
|
||||
return {
|
||||
"type": "state",
|
||||
"observation": self.obs,
|
||||
"enabled": self.enabled,
|
||||
"authorizationGeneration": self.generation,
|
||||
}
|
||||
|
||||
async def push(self):
|
||||
await self.ws.send_json(self.state())
|
||||
|
||||
async def consume(self):
|
||||
async for raw in self.ws:
|
||||
if raw.type != WSMsgType.TEXT:
|
||||
break
|
||||
msg = json.loads(raw.data)
|
||||
if msg["type"] == "stop":
|
||||
self.enabled = False
|
||||
self.obs["paused"] = True
|
||||
self.obs["sequence"] += 1
|
||||
await self.push()
|
||||
elif msg["type"] == "request":
|
||||
self.calls.append(msg)
|
||||
payload, op = msg["payload"], msg["op"]
|
||||
if not self.acknowledge:
|
||||
continue
|
||||
if op == "claim":
|
||||
result = {k: payload[k] for k in ("sessionId", "modelEpoch", "leaseId")}
|
||||
elif op == "action":
|
||||
result = {
|
||||
k: payload[k]
|
||||
for k in ("sessionId", "modelEpoch", "leaseId", "actionSeq", "values")
|
||||
}
|
||||
result["simTime"] = self.obs["simTime"] + 0.002
|
||||
self.obs["simTime"] += 0.002
|
||||
self.obs["sequence"] += 1
|
||||
self.obs["appliedActionSeq"] = payload["actionSeq"]
|
||||
elif op == "reset":
|
||||
self.obs["modelEpoch"] += 1
|
||||
self.obs["sequence"] = 1
|
||||
self.obs["paused"] = True
|
||||
self.obs["simTime"] = 0
|
||||
self.enabled = False
|
||||
result = self.obs
|
||||
else:
|
||||
self.enabled = False
|
||||
self.obs["sequence"] += 1
|
||||
self.obs["paused"] = True
|
||||
result = {"released": True}
|
||||
await self.ws.send_json(
|
||||
{"type": "result", "id": msg["id"], "ok": True, "value": result}
|
||||
)
|
||||
await self.push()
|
||||
|
||||
async def close(self):
|
||||
await self.ws.close()
|
||||
await self.task
|
||||
|
||||
|
||||
class BridgeTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
self.server = TestServer(create_app(TOKEN))
|
||||
await self.server.start_server()
|
||||
self.url = str(self.server.make_url("")).rstrip("/")
|
||||
self.http = ClientSession(headers={"Authorization": f"Bearer {TOKEN}"})
|
||||
self.backend = Backend(self.http, self.url)
|
||||
await self.backend.connect()
|
||||
self.claim_data = {
|
||||
"sessionId": "sim-test",
|
||||
"modelEpoch": 2,
|
||||
"modelFingerprint": FIXTURE["descriptor"]["modelFingerprint"],
|
||||
}
|
||||
|
||||
async def asyncTearDown(self):
|
||||
await self.backend.close()
|
||||
await self.http.close()
|
||||
await self.server.close()
|
||||
|
||||
async def claim(self):
|
||||
response = await self.http.post(self.url + PREFIX + "/lease", json=self.claim_data)
|
||||
self.assertEqual(response.status, 200, await response.text())
|
||||
return await response.json()
|
||||
|
||||
async def test_malformed_tokens_fail_without_encoding_errors(self):
|
||||
for token in ("x" * 15, "a" * 4097, "中文" * 16, "a" * 16 + " "):
|
||||
with self.assertRaises(ValueError):
|
||||
create_app(token)
|
||||
with self.assertRaises(ValueError):
|
||||
SimRobotClient(self.url, token)
|
||||
response = await self.http.get(
|
||||
self.url + PREFIX + "/health", headers={"Authorization": "Bearer " + "é" * 20}
|
||||
)
|
||||
self.assertEqual(response.status, 401)
|
||||
async with self.http.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN) as ws:
|
||||
await ws.send_json({"type": "auth", "protocolVersion": 1, "token": "x" * 16 + "\ud800"})
|
||||
self.assertEqual((await ws.receive_json())["error"]["code"], "UNAUTHORIZED")
|
||||
|
||||
async def test_real_sync_sdk_single_joint_and_reset(self):
|
||||
client = SimRobotClient(self.url, TOKEN)
|
||||
desc = await asyncio.to_thread(client.connect)
|
||||
self.assertEqual(len(desc["actionChannels"]), 1)
|
||||
result = await asyncio.to_thread(client.send_action, {"slider.position": 40})
|
||||
self.assertEqual(result["values"], {"slider.position": 1})
|
||||
obs = await asyncio.to_thread(client.get_observation)
|
||||
self.assertEqual(obs["values"], {"slider.position": 0.12})
|
||||
self.assertEqual(obs["appliedActionSeq"], 1)
|
||||
reset = await asyncio.to_thread(client.reset)
|
||||
self.assertEqual(reset["modelEpoch"], 3)
|
||||
self.assertTrue(reset["paused"])
|
||||
self.assertFalse(client.is_connected)
|
||||
|
||||
async def test_single_browser_single_writer_and_stale_action(self):
|
||||
lease = await self.claim()
|
||||
conflict = await self.http.post(self.url + PREFIX + "/lease", json=self.claim_data)
|
||||
self.assertEqual(conflict.status, 409)
|
||||
other = await self.http.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN)
|
||||
await other.send_json({"type": "auth", "token": TOKEN, "protocolVersion": 1})
|
||||
self.assertEqual((await other.receive_json())["error"]["code"], "CONFLICT")
|
||||
await other.close()
|
||||
packet = {**FIXTURE["validAction"], **lease}
|
||||
headers = {"X-Control-Lease": lease["leaseId"]}
|
||||
good = await self.http.post(self.url + PREFIX + "/action", json=packet, headers=headers)
|
||||
self.assertEqual(good.status, 200)
|
||||
stale = await self.http.post(self.url + PREFIX + "/action", json=packet, headers=headers)
|
||||
self.assertEqual((await stale.json())["error"]["code"], "STALE")
|
||||
self.assertIsNotNone(self.server.app[BROKER].lease)
|
||||
|
||||
async def test_host_origin_auth_and_url_token_rejected(self):
|
||||
for headers in (
|
||||
{"Host": "evil.test"},
|
||||
{"Origin": "http://evil.test"},
|
||||
{"Authorization": "Bearer wrong"},
|
||||
):
|
||||
response = await self.http.get(self.url + PREFIX + "/health", headers=headers)
|
||||
self.assertEqual(response.status, 401)
|
||||
response = await self.http.get(self.url + PREFIX + "/health?token=secret")
|
||||
self.assertEqual(response.status, 400)
|
||||
with self.assertRaises(WSServerHandshakeError):
|
||||
await self.http.ws_connect(self.url + "/ws/control/v1", origin="http://evil.test")
|
||||
ws = await self.http.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN)
|
||||
await ws.send_json({"type": "auth", "token": "wrong", "protocolVersion": 1})
|
||||
self.assertEqual((await ws.receive_json())["error"]["code"], "UNAUTHORIZED")
|
||||
await ws.close()
|
||||
self.assertIsNotNone(self.server.app[BROKER].descriptor)
|
||||
|
||||
async def test_watchdog_and_old_authorization_cannot_reconnect(self):
|
||||
await self.claim()
|
||||
await asyncio.sleep(0.62)
|
||||
self.assertIsNone(self.server.app[BROKER].lease)
|
||||
self.assertFalse(self.backend.enabled)
|
||||
# Replayed enabled heartbeat with the same generation must not reauthorize.
|
||||
self.backend.enabled = True
|
||||
self.backend.obs["paused"] = False
|
||||
self.backend.obs["sequence"] += 1
|
||||
await self.backend.push()
|
||||
response = await self.http.post(self.url + PREFIX + "/lease", json=self.claim_data)
|
||||
self.assertEqual(response.status, 401)
|
||||
self.backend.generation += 1
|
||||
await self.backend.push()
|
||||
await self.claim()
|
||||
|
||||
async def test_old_failed_request_cannot_revoke_a_new_authorization(self):
|
||||
lease = await self.claim()
|
||||
self.backend.acknowledge = False
|
||||
pending = asyncio.create_task(
|
||||
self.http.post(
|
||||
self.url + PREFIX + "/action",
|
||||
json={**FIXTURE["validAction"], **lease},
|
||||
headers={"X-Control-Lease": lease["leaseId"]},
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0.03)
|
||||
self.backend.generation += 1
|
||||
self.backend.obs["sequence"] += 1
|
||||
await self.backend.push()
|
||||
failed = await pending
|
||||
self.assertEqual(failed.status, 503)
|
||||
self.assertTrue(self.server.app[BROKER].authorized)
|
||||
self.assertIsNone(self.server.app[BROKER].lease)
|
||||
self.backend.acknowledge = True
|
||||
new_lease = await self.claim()
|
||||
self.assertNotEqual(new_lease["leaseId"], lease["leaseId"])
|
||||
|
||||
async def test_paused_observation_reports_pause_even_after_freshness_expires(self):
|
||||
broker = self.server.app[BROKER]
|
||||
self.backend.obs["paused"] = True
|
||||
self.backend.obs["sequence"] += 1
|
||||
await broker.state(self.backend.state())
|
||||
for age in (0.0, 1.0):
|
||||
with self.subTest(age=age):
|
||||
broker.observed_at = time.monotonic() - age
|
||||
client = SimRobotClient(self.url, TOKEN)
|
||||
with self.assertRaises(RobotError) as raised:
|
||||
await asyncio.to_thread(client.connect)
|
||||
self.assertEqual(raised.exception.code, "PAUSED")
|
||||
self.assertIn("播放", str(raised.exception))
|
||||
self.assertFalse(client.is_connected)
|
||||
response = await self.http.post(self.url + PREFIX + "/lease", json=self.claim_data)
|
||||
self.assertEqual(response.status, 409)
|
||||
self.assertEqual((await response.json())["error"]["code"], "PAUSED")
|
||||
self.assertIsNone(broker.lease)
|
||||
self.assertFalse(self.backend.calls)
|
||||
|
||||
async def test_stale_observation_is_not_refreshed_by_heartbeat(self):
|
||||
await asyncio.sleep(0.53)
|
||||
await self.backend.push()
|
||||
response = await self.http.get(self.url + PREFIX + "/observation")
|
||||
self.assertEqual((await response.json())["error"]["code"], "STALE")
|
||||
client = SimRobotClient(self.url, TOKEN)
|
||||
with self.assertRaises(RobotError) as raised:
|
||||
await asyncio.to_thread(client.connect)
|
||||
self.assertEqual(raised.exception.code, "STALE")
|
||||
self.assertFalse(client.is_connected)
|
||||
self.assertIsNone(self.server.app[BROKER].lease)
|
||||
|
||||
async def test_rate_limit_and_http_frame_size(self):
|
||||
lease = await self.claim()
|
||||
self.server.app[BROKER].rate = deque([time.monotonic()] * 100)
|
||||
response = await self.http.post(
|
||||
self.url + PREFIX + "/action",
|
||||
json={**FIXTURE["validAction"], **lease},
|
||||
headers={"X-Control-Lease": lease["leaseId"]},
|
||||
)
|
||||
self.assertEqual(response.status, 409)
|
||||
response = await self.http.post(self.url + PREFIX + "/lease", json={"padding": "x" * 70000})
|
||||
self.assertEqual(response.status, 413)
|
||||
|
||||
async def test_no_ack_cancels_pending_and_cannot_keep_lease(self):
|
||||
client = SimRobotClient(self.url, TOKEN)
|
||||
await asyncio.to_thread(client.connect)
|
||||
self.backend.acknowledge = False
|
||||
with self.assertRaises(RobotError):
|
||||
await asyncio.to_thread(client.send_action, {"slider.position": 0.2})
|
||||
self.assertFalse(client.is_connected)
|
||||
self.assertFalse(self.server.app[BROKER].pending)
|
||||
self.assertIsNone(self.server.app[BROKER].lease)
|
||||
|
||||
async def test_latest_request_supersedes_pending(self):
|
||||
lease = await self.claim()
|
||||
self.backend.acknowledge = False
|
||||
headers = {"X-Control-Lease": lease["leaseId"]}
|
||||
one = asyncio.create_task(
|
||||
self.http.post(
|
||||
self.url + PREFIX + "/action",
|
||||
json={**FIXTURE["validAction"], **lease},
|
||||
headers=headers,
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0.03)
|
||||
two = asyncio.create_task(
|
||||
self.http.post(
|
||||
self.url + PREFIX + "/action",
|
||||
json={**FIXTURE["validAction"], **lease, "actionSeq": 2},
|
||||
headers=headers,
|
||||
)
|
||||
)
|
||||
self.assertEqual((await (await one).json())["error"]["code"], "SUPERSEDED")
|
||||
self.assertLessEqual(len(self.server.app[BROKER].pending), 1)
|
||||
self.backend.acknowledge = True
|
||||
# A pending write receives an explicit failure on disconnect, not a false ACK.
|
||||
await self.backend.ws.close()
|
||||
with contextlib.suppress(ConnectionError):
|
||||
response = await two
|
||||
self.assertEqual(response.status, 503)
|
||||
|
||||
async def test_unauthenticated_ws_timeout_and_oversize(self):
|
||||
ws = await self.http.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN)
|
||||
error = await ws.receive_json(timeout=6)
|
||||
self.assertEqual(error["type"], "error")
|
||||
await ws.close()
|
||||
# A too-large peer cannot replace the current backend.
|
||||
ws = await self.http.ws_connect(self.url + "/ws/control/v1", origin=ORIGIN)
|
||||
await ws.send_str("x" * 70000)
|
||||
await ws.receive(timeout=1)
|
||||
await ws.close()
|
||||
self.assertIsNotNone(self.server.app[BROKER].descriptor)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user