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