import { fireEvent, render, screen, waitFor, act } from '@testing-library/react'; import { beforeEach, expect, it, vi } from 'vitest'; import { LocalTrainingPanel } from './LocalTrainingPanel'; import { LocalTrainingClient } from './LocalTrainingClient'; import { MOBILE_TRAINING_TASKS, type MobileDeployment } from '../mobile/training'; import type { TrainingJob, TrainingServerInfo } from './types'; const taskId = 'MobileManipulator-LeKiwi-Bundle'; const deployment = { trainingTaskId: taskId, robotId: 'lekiwi-bundle', browserCompatible: true, trainingStage: 'navigate', evaluation: { episodes: 10, successRate: 0.6, safetyStops: 0, maxJointVelocity: 1.1, meanNavigationDistance: 0.1, seed: 100042, }, } as MobileDeployment; const complete: TrainingJob = { id: 'a'.repeat(32), taskId, state: 'succeeded', createdAt: '', iteration: 2, maxIterations: 2, progress: 1, message: '完成', artifactReady: true, deployment, logs: [ 'Learning iteration 2 / 2', 'Mean value loss: 0.2', 'Mean surrogate loss: -0.1', 'Mean entropy loss: -2', 'Mean reward: 1.25', ], }; const health: TrainingServerInfo = { version: '1', ready: true, trainerRoot: '/training', python: '/python', tasks: ['Unitree-Go2-Flat', ...Object.keys(MOBILE_TRAINING_TASKS)], taskMetadata: Object.entries(MOBILE_TRAINING_TASKS).map(([id, robotId]) => ({ id, robotId, name: id, family: 'mobile-manipulator', browserCompatible: true, terrainPresets: [], terrainParameters: {}, sensorTypes: [], sensorParameters: {}, mapSyncScope: '', })), }; beforeEach(() => { vi.restoreAllMocks(); localStorage.clear(); sessionStorage.clear(); vi.spyOn(LocalTrainingClient.prototype, 'health').mockResolvedValue(health); vi.spyOn(LocalTrainingClient.prototype, 'presets').mockResolvedValue([]); }); async function connect() { fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'test' } }); fireEvent.click(screen.getByRole('button', { name: '连接' })); await screen.findByRole('button', { name: '发起本地训练' }); } it('自动选择变体、隐藏Go2配置;上传快照→创建作业→下载元数据与ONNX→主会话导入', async () => { const snapshot = new File(['scene'], 'scene.zip'); const policy = new File(['onnx'], 'policy.onnx'); const bridge = { robotId: 'lekiwi-bundle', prepare: vi.fn().mockResolvedValue(snapshot), importPolicy: vi.fn().mockResolvedValue(undefined), }; const upload = vi .spyOn(LocalTrainingClient.prototype, 'uploadMobileScene') .mockResolvedValue({ id: 'b'.repeat(64), robotId: 'lekiwi-bundle', sceneSha256: 's' }); const start = vi.spyOn(LocalTrainingClient.prototype, 'startJob').mockResolvedValue(complete); vi.spyOn(LocalTrainingClient.prototype, 'downloadPolicy').mockResolvedValue(policy); vi.spyOn(LocalTrainingClient.prototype, 'downloadMobileDeployment').mockResolvedValue(deployment); const go2 = vi.fn(); render(); await connect(); expect(screen.getByLabelText('训练任务')).toHaveValue(taskId); expect(screen.queryByLabelText('训练地形')).not.toBeInTheDocument(); expect(screen.queryByLabelText('基础策略')).not.toBeInTheDocument(); expect(screen.queryByLabelText('W&B 模式')).not.toBeInTheDocument(); expect(screen.getByLabelText('并行环境')).toHaveAttribute('max', '64'); expect(screen.getByLabelText('计算设备')).toHaveValue('cpu'); fireEvent.change(screen.getByLabelText('每环境采样步数'), { target: { value: '8' } }); fireEvent.change(screen.getByLabelText('训练迭代'), { target: { value: '2' } }); fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: '0.7' } }); fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); await waitFor(() => expect(start).toHaveBeenCalledOnce()); expect(bridge.prepare).toHaveBeenCalledWith(taskId); expect(upload).toHaveBeenCalledWith(snapshot); expect(start.mock.calls[0][0]).toMatchObject({ taskId, numEnvs: 1, maxIterations: 2, device: 'cpu', seed: 42, mobilePackageId: 'b'.repeat(64), mobileParams: { rolloutSteps: 8, goalPosition: [0.7, 0.15, 0.019], stage: 'navigate', evaluationEpisodes: 10, positionJitter: 0.1, navigationBootstrapSteps: 4096, }, }); expect(start.mock.calls[0][0]).not.toHaveProperty('terrainPreset'); expect(await screen.findByText('价值损失')).toBeInTheDocument(); fireEvent.click(screen.getByRole('button', { name: '导入策略' })); await waitFor(() => expect(bridge.importPolicy).toHaveBeenCalledWith(policy, deployment)); expect(go2).not.toHaveBeenCalled(); expect(screen.getByText('60.0% / 10 回合')).toBeVisible(); expect(screen.getByText('1.100 rad/s')).toBeVisible(); fireEvent.click(screen.getByRole('button', { name: '接续此作业(保留权重)' })); expect(screen.getByLabelText('接续作业 ID')).toHaveValue(complete.id); fireEvent.change(screen.getByLabelText('移动操作训练阶段'), { target: { value: 'reach' } }); start.mockRejectedValueOnce(new Error('上一阶段尚未达标')); fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); expect(await screen.findByRole('alert')).toHaveTextContent('上一阶段尚未达标'); expect(start.mock.calls[1][0].mobileParams).toMatchObject({ stage: 'reach', sourceJobId: complete.id, }); }); it('场景准备失败不创建作业;任务切换恢复Go2参数', async () => { const start = vi.spyOn(LocalTrainingClient.prototype, 'startJob'); render( , ); await connect(); fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); expect(await screen.findByRole('alert')).toHaveTextContent('机器人绑定失败'); expect(start).not.toHaveBeenCalled(); fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } }); expect(screen.queryByLabelText('每环境采样步数')).not.toBeInTheDocument(); expect(screen.getByLabelText('并行环境')).toHaveValue(4096); expect(screen.getByLabelText('计算设备')).toHaveValue('gpu'); }); it('轮询移动作业日志和完成状态;拒绝其他任务元数据', async () => { const bridge = { robotId: 'lekiwi-bundle', prepare: vi.fn(), importPolicy: vi.fn() }; vi.spyOn(LocalTrainingClient.prototype, 'health').mockResolvedValue({ ...health, activeJobId: complete.id, }); vi.spyOn(LocalTrainingClient.prototype, 'job') .mockResolvedValueOnce({ ...complete, state: 'running', artifactReady: false }) .mockResolvedValue(complete); vi.spyOn(LocalTrainingClient.prototype, 'downloadMobileDeployment').mockResolvedValue({ ...deployment, trainingTaskId: 'MobileManipulator-LeKiwi-v1', }); render(); fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'test' } }); fireEvent.click(screen.getByRole('button', { name: '连接' })); await screen.findByRole('button', { name: '停止训练' }); // Trigger the existing 1500ms status polling path without changing production timers. await act(async () => { await new Promise((resolve) => setTimeout(resolve, 1600)); }); fireEvent.click(await screen.findByRole('button', { name: '导入策略' })); expect(await screen.findByRole('alert')).toHaveTextContent('成果物与机器人变体不匹配'); expect(bridge.importPolicy).not.toHaveBeenCalled(); });