3ad29356c9
集成通用机器人数值接口、本机控制桥、LeRobot 插件和统一键盘遥操作。采用离线 CoACD 全臂碰撞配方 revision 4、局部装配区切分与结构自接触,限制直接关节位姿写入并保留安全看门狗。同步版本号、变更记录、来源许可证和兼容性验证。
323 lines
14 KiB
Python
323 lines
14 KiB
Python
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()
|