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()