import { readFileSync } from 'node:fs'; import { beforeEach, describe, expect, it, vi } from 'vitest'; import type { ProjectManifest } from '../project/types'; import fixture from '../rl/fixtures/obstacleDeployment.json'; import { validatePolicyDeployment } from '../rl/deployment'; const ort = vi.hoisted(() => ({ create: vi.fn() })); vi.mock('onnxruntime-web/wasm', () => ({ env: { wasm: {} }, InferenceSession: { create: ort.create }, })); vi.mock('@mujoco/mujoco', async (original) => { const actual = await original(); return { ...actual, default: () => actual.default({ wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm') }), }; }); import { MainThreadPhysicsAdapter } from './PhysicsAdapter'; const deployment = validatePolicyDeployment(fixture); function graph(size = 47) { return { inputNames: ['obs'], outputNames: ['actions'], inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }], outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }], release: vi.fn().mockResolvedValue(undefined), }; } function project(): ProjectManifest { const doc = new DOMParser().parseFromString( readFileSync('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', 'utf8'), 'application/xml', ); doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove()); const actuators = doc.createElement('actuator'); for (const name of deployment.jointNames) { const motor = doc.createElement('motor'); motor.setAttribute('joint', name); motor.setAttribute('name', `${name}_motor`); actuators.append(motor); } doc.documentElement.append(actuators); const data = new TextEncoder().encode(new XMLSerializer().serializeToString(doc)); return { id: 'transaction', name: 'go2', files: [{ path: 'go2.xml', data, size: data.length, source: 'file', mimeType: 'text/xml' }], entries: [{ path: 'go2.xml', format: 'mjcf', label: 'go2' }], maps: [], selectedEntry: 'go2.xml', totalBytes: data.length, }; } async function existing() { const adapter = new MainThreadPhysicsAdapter(), source = project(), oldGraph = graph(); await adapter.load(source, 'go2.xml'); ort.create.mockResolvedValueOnce(oldGraph); await adapter.loadRLPolicy(new Uint8Array([1]), 'old-flat.onnx'); adapter.setRLPolicyEnabled(true); const old = adapter.session!; old.data.time = 7; old.data.qvel[0] = 0.123; old.data.qpos[0] = 1.25; old.data.ctrl[0] = 0.3; adapter.setPaused(false); return { adapter, source, old, oldGraph, before: adapter.snapshot()! }; } beforeEach(() => { vi.clearAllMocks(); ort.create.mockReset(); }); describe('PhysicsAdapter training transaction', () => { it.each(['wrong graph', 'ORT initialization', 'binding'] as const)( '%s失败保留原session/workspace/策略及物理状态', async (failure) => { const { adapter, source, old, oldGraph, before } = await existing(); const workspace = adapter.workspace, dispose = vi.spyOn(old, 'dispose'); const candidate = graph(47); if (failure === 'ORT initialization') ort.create.mockRejectedValueOnce(new Error('ORT init failed')); else ort.create.mockResolvedValueOnce(candidate); if (failure === 'binding') { const text = new TextDecoder() .decode(source.files[0].data) .replace('name="FL_hip_joint_motor"', 'name="incompatible_name"'); source.files[0].data = new TextEncoder().encode(text); } try { await expect( adapter.load(source, 'go2.xml', { trainingDeployment: deployment, trainingPolicy: { data: new Uint8Array([2]), path: 'candidate.onnx' }, }), ).rejects.toThrow( failure === 'wrong graph' ? /维度/ : failure === 'binding' ? /驱动器/ : /ORT init/, ); expect(adapter.session).toBe(old); expect(adapter.workspace).toBe(workspace); expect(adapter.snapshot()).toEqual(before); expect(dispose).not.toHaveBeenCalled(); expect(oldGraph.release).not.toHaveBeenCalled(); if (failure === 'wrong graph') expect(candidate.release).toHaveBeenCalledOnce(); } finally { adapter.dispose(); } }, ); it('完成ORT初始化前不发布候选,成功后仍保留旧资源以供viewer失败回滚', async () => { const { adapter, source, old, oldGraph, before } = await existing(); let finish!: (value: ReturnType) => void; let started!: () => void; const initialized = new Promise((resolve) => { started = resolve; }); ort.create.mockImplementationOnce(() => { started(); return new Promise((resolve) => { finish = resolve; }); }); const candidate = graph(81); try { const loading = adapter.load(source, 'go2.xml', { trainingDeployment: deployment, trainingPolicy: { data: new Uint8Array([2]), path: 'new.onnx' }, }); await initialized; expect(adapter.session).toBe(old); expect(oldGraph.release).not.toHaveBeenCalled(); finish(candidate); await loading; expect(adapter.session).not.toBe(old); expect(adapter.snapshot()?.rlPolicy).toMatchObject({ path: 'new.onnx', observationSize: 81, enabled: true, }); expect(oldGraph.release).not.toHaveBeenCalled(); adapter.rollbackRetired(); expect(adapter.session).toBe(old); expect(adapter.snapshot()).toEqual(before); for (let i = 0; i < 10; i++) await Promise.resolve(); expect(candidate.release).toHaveBeenCalledOnce(); expect(oldGraph.release).not.toHaveBeenCalled(); } finally { adapter.dispose(); } }); }); it('候选ORT等待期间另一场景完成加载,迟到候选不得覆盖最新场景', async () => { const { adapter, source } = await existing(); let finish!: (value: ReturnType) => void; let started!: () => void; const initialized = new Promise((resolve) => { started = resolve; }); ort.create.mockImplementationOnce(() => { started(); return new Promise((resolve) => { finish = resolve; }); }); const candidate = graph(81); try { const stale = adapter.load(source, 'go2.xml', { trainingDeployment: deployment, trainingPolicy: { data: new Uint8Array([2]), path: 'stale.onnx' }, }); await initialized; await adapter.load(source, 'go2.xml'); const latest = adapter.session, before = adapter.snapshot(); const rejected = expect(stale).rejects.toThrow(/取消/); finish(candidate); await rejected; expect(adapter.session).toBe(latest); expect(adapter.snapshot()).toEqual(before); for (let i = 0; i < 10; i++) await Promise.resolve(); expect(candidate.release).toHaveBeenCalledOnce(); } finally { adapter.dispose(); } }); describe('PhysicsAdapter custom_boxes applied scene compiler', () => { it('真实WASM读取工程路径/嵌套body变换/旋转多实例;与共享CPU布局halfsize和原点一致', async () => { const { default: shared } = await import('../../../training_server/tests/fixtures/custom-boxes.json'); const source = project(); const descriptor = { schemaVersion: 1, id: 'warehouse', name: '仓库', coordinateSystem: { units: 'm', up: 'Z', forward: '+X' }, physics: { source: 'physics/world.xml' }, spawnPoints: [], }; for (const [path, text] of [ ['maps/warehouse/map.json', JSON.stringify(descriptor)], [ 'maps/warehouse/physics/world.xml', '', ], ]) { const data = new TextEncoder().encode(text); source.files.push({ path, data, size: data.length, source: 'file', mimeType: 'text/xml' }); } const assets = [ { id: 'one', name: 'one', selection: { kind: 'project' as const, descriptorPath: 'maps/warehouse/map.json', positionX: 1, positionY: 1, }, }, ]; const adapter = new MainThreadPhysicsAdapter(); try { await adapter.load(source, 'go2.xml', { mapAssets: assets }); const initial = adapter.exportTrainingTerrain(assets, { spawn: [-2, -1], target: [2, -1] }); // Geometry from actual WASM matches the JSON consumed by CPU TerrainGenerator tests. expect(initial.boxes[1]).toEqual(shared.boxes[1]); expect(initial.actualObstacleCount).toBe(1); adapter.session!.data.qpos[0] = 8; expect(adapter.exportTrainingTerrain(assets).spawn[0]).not.toBe(8); expect(() => adapter.exportTrainingTerrain([])).toThrow(/过时/); const multi = [ ...assets, { id: 'two', name: 'two', selection: { ...assets[0].selection, positionX: -2, positionY: 1, yawDeg: 45 }, }, ]; await adapter.load(source, 'go2.xml', { mapAssets: multi }); const layout = adapter.exportTrainingTerrain(multi, { spawn: [-2, -1], target: [2, -1] }); expect(layout.boxes).toHaveLength(3); expect(layout.boxes[2].pos[0]).toBeCloseTo(-2 - Math.SQRT1_2, 12); expect(layout.boxes[2].pos[1]).toBeCloseTo(1 + Math.SQRT1_2, 12); expect(layout.boxes[2].size[0]).toBeCloseTo(0.7 * Math.SQRT1_2, 12); expect(layout.boxes[2].size[1]).toBeCloseTo(0.7 * Math.SQRT1_2, 12); adapter.rollbackRetired(); expect(adapter.exportTrainingTerrain(assets).actualObstacleCount).toBe(1); // A deployment replacement cannot pretend to be the original applied editor map. await adapter.load(source, 'go2.xml', { mapAssets: assets, trainingDeployment: deployment }); expect(() => adapter.exportTrainingTerrain(assets)).toThrow(/已应用/); } finally { adapter.dispose(); } }); });