import { beforeEach, describe, expect, it, vi } from 'vitest'; import fixture from '../fixtures/obstacleDeployment.json'; import { validatePolicyDeployment } from '../deployment'; const mock = vi.hoisted(() => ({ create: vi.fn(), tensors: [] as { dispose: ReturnType }[], })); vi.mock('onnxruntime-web/wasm', () => ({ env: { wasm: {} }, InferenceSession: { create: mock.create }, Tensor: class { dispose = vi.fn(); constructor( public type: string, public data: Float32Array, public dims: number[], ) { mock.tensors.push(this); } }, })); import { OnnxPolicyRuntime } from './OnnxPolicyRuntime'; const deployment = validatePolicyDeployment(fixture); const flush = async () => { for (let i = 0; i < 10; i++) await Promise.resolve(); }; function setup(size = 81) { const session = { inputNames: ['obs'], outputNames: ['actions'], inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }], outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }], run: vi.fn(), release: vi.fn().mockResolvedValue(undefined), }; const bindings = { observe: vi.fn(() => new Float32Array(size)), apply: vi.fn(), clear: vi.fn(), reset: vi.fn(), }; mock.create.mockResolvedValue(session); return { session, bindings }; } beforeEach(() => { vi.clearAllMocks(); mock.tensors.length = 0; }); describe('OnnxPolicyRuntime held-action', () => { it('动态目标转发不重建ORT、不启用策略或打断in-flight,status实时读取', async () => { const { session, bindings } = setup(); let finish!: (outputs: Record) => void; session.run.mockImplementation( () => new Promise((resolve) => { finish = resolve; }), ); const navigation = { target: [1, 2] as [number, number], defaultTarget: [5, 0] as [number, number], distance: 3, targetHeight: 0, }; const targetBindings = { ...bindings, navigationStatus: () => navigation, setNavigationTarget: vi.fn(), resetNavigationTarget: vi.fn(), }; const runtime = await OnnxPolicyRuntime.load( new Uint8Array(), 'goal.onnx', targetBindings, deployment, ); runtime.setNavigationTarget([2, 3]); expect(targetBindings.setNavigationTarget).toHaveBeenCalledWith([2, 3]); expect(runtime.status().enabled).toBe(false); expect(runtime.status().navigation).toBe(navigation); runtime.setEnabled(true, 0); runtime.step(0); runtime.setNavigationTarget([3, 4]); runtime.step(0.02); expect(session.run).toHaveBeenCalledOnce(); expect(mock.create).toHaveBeenCalledOnce(); runtime.resetNavigationTarget(); expect(targetBindings.resetNavigationTarget).toHaveBeenCalledOnce(); finish({ actions: { type: 'float32', data: new Float32Array(12), dispose: vi.fn() } }); await flush(); runtime.dispose(); expect(runtime.status().navigation).toBeUndefined(); }); it('81维契约单in-flight,持有最新动作,不阻塞物理步;reset拒绝旧promise', async () => { const { session, bindings } = setup(); let finish!: (outputs: Record) => void; session.run.mockImplementation( () => new Promise((resolve) => { finish = resolve; }), ); const runtime = await OnnxPolicyRuntime.load( new Uint8Array([1]), 'policy.onnx', bindings, deployment, ); runtime.setEnabled(true, 0); runtime.step(0); runtime.step(0.02); runtime.step(0.04); expect(session.run).toHaveBeenCalledOnce(); expect(bindings.apply).toHaveBeenCalledTimes(3); const output = { type: 'float32', data: new Float32Array(12).fill(0.3), dispose: vi.fn() }; finish({ actions: output }); await flush(); runtime.step(0.06); expect(bindings.apply.mock.lastCall?.[0][0]).toBeCloseTo(0.3); expect(session.run).toHaveBeenCalledTimes(2); expect(output.dispose).toHaveBeenCalledOnce(); runtime.reset(0); finish({ actions: { ...output, data: new Float32Array(12).fill(0.9) } }); await flush(); runtime.step(0); expect(bindings.apply.mock.lastCall?.[0][0]).toBe(0); expect(bindings.reset).toHaveBeenCalledWith(0); runtime.dispose(); expect(session.release).not.toHaveBeenCalled(); finish({ actions: output }); await flush(); expect(session.release).toHaveBeenCalledOnce(); expect(mock.tensors.every((t) => t.dispose.mock.calls.length === 1)).toBe(true); }); it('81策略不能默认为47、Rough维度不能加载,失败释放session', async () => { const { session, bindings } = setup(); await expect(OnnxPolicyRuntime.load(new Uint8Array(), 'wrong.onnx', bindings)).rejects.toThrow( /维度/, ); expect(session.release).toHaveBeenCalledOnce(); session.inputMetadata[0].shape = [1, 234]; await expect( OnnxPolicyRuntime.load(new Uint8Array(), 'rough.onnx', bindings, deployment), ).rejects.toThrow(/维度/); }); it('原47维Flat保持向后兼容,非法观测/输出失败关闭并清动作', async () => { const { session, bindings } = setup(47); session.run.mockResolvedValue({ actions: { type: 'float32', data: new Float32Array(12).fill(NaN), dispose: vi.fn() }, }); const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'flat.onnx', bindings); expect(runtime.status().observationSize).toBe(47); runtime.setEnabled(true, 0); runtime.step(0); await flush(); expect(runtime.status().enabled).toBe(false); expect(runtime.status().error).toMatch(/非有限/); runtime.reset(0); bindings.observe.mockReturnValue(new Float32Array(81)); runtime.setEnabled(true, 0); runtime.step(0); expect(runtime.status().error).toMatch(/维度/); runtime.dispose(); await flush(); }); }); it('默认Flat作业的兼容契约仍检查真实graph,拒绝81维及动态特征', async () => { const flat = validatePolicyDeployment({ ...fixture, taskId: 'Unitree-Go2-Flat', observationSize: 47, observationTerms: fixture.observationTerms.slice(0, 7), terrain: undefined, terrainPreset: undefined, terrainParams: undefined, sensorCfg: undefined, navigation: undefined, }); const { session, bindings } = setup(47); const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'legacy.onnx', bindings, flat); expect(runtime.status().observationSize).toBe(47); runtime.dispose(); await flush(); session.inputMetadata[0].shape = [1, 81]; await expect( OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-wrong.onnx', bindings, flat), ).rejects.toThrow(/维度/); session.inputMetadata[0].shape = [1, -1]; await expect( OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-dynamic.onnx', bindings, flat), ).rejects.toThrow(/固定/); }); it('97 metadata不能冒充81 graph;真实97 shape沿用single-flight生命周期', async () => { const multi = validatePolicyDeployment({ ...fixture, observationSize: 97, sensorCfg: { ...fixture.sensorCfg, sensorMode: 'multi_ring_raycast', rayCount: 48 }, }); const { session, bindings } = setup(81); await expect( OnnxPolicyRuntime.load(new Uint8Array([1]), 'multi.onnx', bindings, multi), ).rejects.toThrow(/维度/); expect(session.release).toHaveBeenCalledOnce(); const next = setup(97); const runtime = await OnnxPolicyRuntime.load( new Uint8Array([1]), 'multi.onnx', next.bindings, multi, ); expect(runtime.status().observationSize).toBe(97); runtime.dispose(); await flush(); expect(next.session.release).toHaveBeenCalledOnce(); });